Compare commits
@@ -65,7 +65,7 @@ def test_cl_combinatorial(self):
|
||||
{} # 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 |
|
||||
| `ppp` | `PromptPostProcessor \| str \| None` | Supported values `"nocup"`, `"nostrict"` or a specific instance |
|
||||
| `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
|
||||
|
||||
@@ -94,23 +97,35 @@ Do not use bare `assert` statements.
|
||||
|
||||
## 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
|
||||
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"""
|
||||
self.process(
|
||||
InputTuple("a, , b", ""),
|
||||
OutputTuple("a | b", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.def_env_info,
|
||||
replace(
|
||||
self.defopts,
|
||||
keep_choices_order=True,
|
||||
cup_do_cleanup=False,
|
||||
do_combinatorial=True,
|
||||
self.def_env_info,
|
||||
model_filename="./webui/models/Stable-diffusion/testmodel.safetensors",
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
self.wildcards_obj,
|
||||
|
||||
@@ -8,3 +8,5 @@ logs
|
||||
tests/tests_local.py
|
||||
tests/local_wildcards
|
||||
tests/logs
|
||||
tools/*.bat
|
||||
dev/*.bat
|
||||
|
||||
@@ -85,11 +85,18 @@ See the [cookbook](docs/COOKBOOK.md) for interesting usages.
|
||||
|
||||
## 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
|
||||
|
||||
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
|
||||
|
||||
|
||||
@@ -10,8 +10,10 @@ from pathlib import Path
|
||||
|
||||
sys.path.append(str(Path(__file__).resolve().parent))
|
||||
|
||||
# pylint: disable=wrong-import-position
|
||||
from .ppp_comfyui import (
|
||||
PromptPostProcessorComfyUINode,
|
||||
PromptPostProcessorRunModeOptionsComfyUINode,
|
||||
PromptPostProcessorWildcardOptionsComfyUINode,
|
||||
PromptPostProcessorENMappingOptionsComfyUINode,
|
||||
PromptPostProcessorSTNOptionsComfyUINode,
|
||||
@@ -22,6 +24,7 @@ from .ppp_comfyui import (
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"ACBPromptPostProcessor": PromptPostProcessorComfyUINode,
|
||||
"ACBPPPRunModeOptions": PromptPostProcessorRunModeOptionsComfyUINode,
|
||||
"ACBPPPWildcardOptions": PromptPostProcessorWildcardOptionsComfyUINode,
|
||||
"ACBPPPENMappingOptions": PromptPostProcessorENMappingOptionsComfyUINode,
|
||||
"ACBPPPSendToNegativeOptions": PromptPostProcessorSTNOptionsComfyUINode,
|
||||
@@ -31,6 +34,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ACBPromptPostProcessor": "ACB Prompt Post Processor",
|
||||
"ACBPPPRunModeOptions": "ACB PPP Run Mode Options",
|
||||
"ACBPPPWildcardOptions": "ACB PPP Wildcard Options",
|
||||
"ACBPPPENMappingOptions": "ACB PPP ExtraNetwork Mapping Options",
|
||||
"ACBPPPSendToNegativeOptions": "ACB PPP Send-To-Negative Options",
|
||||
|
||||
+48
-9
@@ -12,7 +12,7 @@ The model variants now support regular expressions instead of a list of strings
|
||||
|
||||
## 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.
|
||||
|
||||
@@ -35,14 +35,13 @@ Inputs:
|
||||
* **process_wildcards**: Activates the wildcard processing.
|
||||
* **do_cleanup**: Activates the cleanup processing.
|
||||
* **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.
|
||||
* **combinatorial_shuffle**: It shuffles the combinatorial results.
|
||||
* **combinatorial_limit**: Limit for the number of generated combinations.
|
||||
* **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.
|
||||
* **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.
|
||||
* **wc_options**: Connection to a Wildcards options node.
|
||||
* **stn_options**: Connection to a Send-To-Negative options node.
|
||||
* **cup_options**: Connection to a Cleanup 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.
|
||||
|
||||
@@ -58,7 +57,19 @@ Outputs:
|
||||
* **neg_prompt**: Resulting negative prompt.
|
||||
* **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
|
||||
|
||||
@@ -94,6 +105,20 @@ Output:
|
||||
|
||||
* **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
|
||||
|
||||
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.
|
||||
* **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.
|
||||
* **Combinatorial mode**: Generate all possible prompt combinations (from choices and wildcards) and cycle through them to fill the batch.
|
||||
* **Shuffle combinations**: It shuffles the combinatorial results.
|
||||
* **Combinations limit**: Maximum number of combinations to generate (0 = no limit). The actual maximum limit is the number of images (batch size * count).
|
||||
* **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.
|
||||
* **Results 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
|
||||
|
||||
|
||||
+52
-6
@@ -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.
|
||||
|
||||
## 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
|
||||
|
||||
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:
|
||||
|
||||
| Extension | Format |
|
||||
|------------------|----------------------------------------------------|
|
||||
| `.yaml` / `.yml` | YAML list of records |
|
||||
| `.jsonl` | JSON Lines, one JSON object per line |
|
||||
| `.csv` | CSV with a header row (semicolon-delimited) |
|
||||
| anything else | Plain text with labelled sections |
|
||||
| Extension | Format |
|
||||
|------------------|---------------------------------------------|
|
||||
| `.yaml` / `.yml` | YAML list of records |
|
||||
| `.jsonl` | JSON Lines, one JSON object per line |
|
||||
| `.csv` | CSV with a header row (semicolon-delimited) |
|
||||
| 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).
|
||||
|
||||
|
||||
+53
-41
@@ -1,5 +1,9 @@
|
||||
# Prompt PostProcessor syntax
|
||||
|
||||
## Basic usage
|
||||
|
||||
The extension works by modifying the prompt and negative prompt, expanding, replacing and cleaning content, before the image is generated.
|
||||
|
||||
## Commands
|
||||
|
||||
The extension uses a format for its commands similar to an extranetwork, but it has a `ppp:` prefix followed by the command, and then a space and any parameters (if any).
|
||||
@@ -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):
|
||||
|
||||
* `~` (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.
|
||||
* `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.
|
||||
@@ -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.
|
||||
* `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.
|
||||
* `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)
|
||||
|
||||
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:
|
||||
|
||||
| Construct | Result |
|
||||
| --------- | ------ |
|
||||
|---------------------------------------------------|---------------------------------------------------------------------------------------------------------------------------------------|
|
||||
| `{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 |
|
||||
| `{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:
|
||||
|
||||
| Construct | Result |
|
||||
| --------- | ------ |
|
||||
|--------------------------------------|--------------------------------------------------------------------------|
|
||||
| `__wildcard__` | select 1 choice |
|
||||
| `__path/wildcard'0'__` | select the first 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`: 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)
|
||||
* `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:
|
||||
|
||||
| Construct | Meaning |
|
||||
| --------- | ------- |
|
||||
|-----------------------------------------------|--------------------|
|
||||
| `<ppp:setwcdeffilter 'identifier' 'filter'/>` | Sets a 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:
|
||||
|
||||
| System variable | Value |
|
||||
| --------------- | ----- |
|
||||
| `_model` | the model identifier (`sd1`, `sd2`, `sdxl`, `sd3`, `flux`, `auraflow`). `_sd` also works but is deprecated. |
|
||||
| `_modelname` | the model filename (without path). Do not confuse with the `modelname` input in *ComfyUI* which matches actually to the `_modelfullname` variable. `_sdname` also works but is deprecated. |
|
||||
| `_modelfullname` | the model filename (with path). `_sdfullname` also works but is deprecated. In *ComfyUI* this variable can also be **set** to override the filename used for model detection (see below). |
|
||||
| `_modelclass` | the class used for the model. Note that this is dependent on the webui. In A1111 all SD versions use the same class. Can be used for new models that are not supported yet with the `_is_*` variables. The debug setting will show all system variables when generating in case you need to see which one to use for a certain model. |
|
||||
| `_is_kkkk` | true if the model is of kind *kkkk* (the model identifier, f.e. sdxl; those set in the ppp_config.yaml file) |
|
||||
| `_is_vvvv` | true if the model matches the *vvvv* model variant definition (based on its filename). Note that the corresponding variable for the model kind will also be true. |
|
||||
| `_is_pure_kkkk` | true if the model is of kind *kkkk* and not a variant. |
|
||||
| `_is_variant_kkkk` | true if the model version is any variant of model kind *kkkk* and not the pure version. Note that the corresponding variable for the model kind will also be true. |
|
||||
| `_is_sd` | true if the model is any version of SD |
|
||||
| `_is_ssd` | true if the model is SSD (Segmind Stable Diffusion 1B). Note that for an SSD model `_is_sdxl` will also be true. |
|
||||
| `_is_sdxl_no_ssd` | true if the model is SDXL and not an SSD model. |
|
||||
| `_is_sdxl_no_pony` | true if the model is SDXL and not a Pony model (the `pony` variant must be defined in settings). Kept to maintain compatibility with previous versions. |
|
||||
| `_opt_...` | All the options. |
|
||||
| `_input_seed` | The seed used. |
|
||||
| `_input_pos_prompt` | The original positive prompt. |
|
||||
| `_input_neg_prompt` | The original negative prompt. |
|
||||
| System variable | Value |
|
||||
|--------------------------|---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
|
||||
| `_model` | The model identifier (`sd1`, `sd2`, `sdxl`, `sd3`, `flux`, `auraflow`). `_sd` also works but is deprecated. |
|
||||
| `_modelname` | The model filename (without path). Do not confuse with the `modelname` input in *ComfyUI* which matches actually to the `_modelfullname` variable. `_sdname` also works but is deprecated. |
|
||||
| `_modelfullname` | The model filename (with path). `_sdfullname` also works but is deprecated. In *ComfyUI* this variable can also be **set** to override the filename used for model detection (see below). |
|
||||
| `_modelclass` | The class used for the model. Note that this is dependent on the webui. In A1111 all SD versions use the same class. Can be used for new models that are not supported yet with the `_is_*` variables. The debug setting will show all system variables when generating in case you need to see which one to use for a certain model. |
|
||||
| `_is_kkkk` | True if the model is of kind *kkkk* (the model identifier, f.e. sdxl; those set in the ppp_config.yaml file) |
|
||||
| `_is_vvvv` | True if the model matches the *vvvv* model variant definition (based on its filename). Note that the corresponding variable for the model kind will also be true. |
|
||||
| `_is_pure_kkkk` | True if the model is of kind *kkkk* and not a variant. |
|
||||
| `_is_variant_kkkk` | True if the model version is any variant of model kind *kkkk* and not the pure version. Note that the corresponding variable for the model kind will also be true. |
|
||||
| `_is_sd` | True if the model is any version of SD |
|
||||
| `_is_ssd` | True if the model is SSD (Segmind Stable Diffusion 1B). Note that for an SSD model `_is_sdxl` will also be true. |
|
||||
| `_is_sdxl_no_ssd` | True if the model is SDXL and not an SSD model. |
|
||||
| `_is_sdxl_no_pony` | True if the model is SDXL and not a Pony model (the `pony` variant must be defined in settings). Kept to maintain compatibility with previous versions. |
|
||||
| `_opt_...` | All the options. |
|
||||
| `_input_seed` | The starting seed used. |
|
||||
| `_input_pos_prompt` | The original positive prompt. |
|
||||
| `_input_neg_prompt` | The original negative prompt. |
|
||||
| `_input_prev_pos_prompt` | The positive prompt result of the previous phase. Available only in the hires fix phase in A1111 compatible hosts and with no combinatorial generation. |
|
||||
| `_input_prev_neg_prompt` | The negative prompt result of the previous phase. Available only in the hires fix phase in A1111 compatible hosts and with no combinatorial generation. |
|
||||
| `_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]
|
||||
> The model path is relative to the checkpoint/difussion_models folder, just as it appears in the load nodes.
|
||||
|
||||
> [!NOTE]
|
||||
> You can use the `_input_prev_pos_prompt` and `_input_prev_neg_prompt` only in the hires fix prompt boxes in A1111 compatible hosts, and only without the combinatorial option.
|
||||
|
||||
## Set command
|
||||
|
||||
This command sets the value of a variable that can be checked later.
|
||||
@@ -249,14 +261,14 @@ The `add` and `ifundefined` modifiers are mutually exclusive and cannot be used
|
||||
The *Dynamic Prompts* format also works:
|
||||
|
||||
| Construct | Meaning |
|
||||
| --------- | ------- |
|
||||
|-----------------|----------------------|
|
||||
| `${var=value}` | regular evaluation |
|
||||
| `${var=!value}` | immediate evaluation |
|
||||
|
||||
If also supports the addition and undefined check as an extension of the *Dynamic Prompts* format:
|
||||
|
||||
| Construct | Meaning |
|
||||
| --------- | ------- |
|
||||
|------------------|--------------------------------------|
|
||||
| `${var+=value}` | equivalent to `add` |
|
||||
| `${var+=!value}` | equivalent to `evaluate add` |
|
||||
| `${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:
|
||||
|
||||
| Construct |
|
||||
| --------- |
|
||||
|----------------------------------------|
|
||||
| `<ppp:echo varname/>` |
|
||||
| `<ppp:echo varname>default<ppp:/echo>` |
|
||||
|
||||
The *Dynamic Prompts* format is:
|
||||
|
||||
| Construct |
|
||||
| --------- |
|
||||
|----------------------|
|
||||
| `${varname}` |
|
||||
| `${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:
|
||||
|
||||
| Construct | Meaning |
|
||||
| --------- | ------- |
|
||||
|----------------------------|------------------------------------------------------------------|
|
||||
| `${var[]=value}` | initialize and set the first value |
|
||||
| `${var[]=*()}` | initialize an empty 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
|
||||
* An ampersand followed by a string inside the brackets (with quotes) is used to get the full array joined with a separator.
|
||||
|
||||
| Construct | Meaning |
|
||||
| --------- | ------- |
|
||||
| `${var[]}` | echo all elements with a default separator |
|
||||
| `${var[&' / ']}` | echo all elements with a specific separator |
|
||||
| `${var[n]}` | echo an element from the array |
|
||||
| `${var[n]:default}` | echo an element with a default |
|
||||
| `${var[#]}` | echo the length of the array |
|
||||
| Construct | Meaning |
|
||||
|---------------------|---------------------------------------------|
|
||||
| `${var[]}` | echo all elements with a default separator |
|
||||
| `${var[&' / ']}` | echo all elements with a specific separator |
|
||||
| `${var[n]}` | echo an element from the array |
|
||||
| `${var[n]:default}` | echo an element with a default |
|
||||
| `${var[#]}` | echo the length of the array |
|
||||
|
||||
## 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:
|
||||
|
||||
| Construct | Meaning |
|
||||
| --------- | ------- |
|
||||
|-------------------------------------|---------------------------------------------------------------------------------|
|
||||
| `operand` | check truthyness of the operand, meaning not zero, empty string nor empty array |
|
||||
| `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).
|
||||
|
||||
| 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 |
|
||||
| `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 |
|
||||
@@ -458,10 +470,10 @@ extnettype:
|
||||
|
||||
Used like this:
|
||||
|
||||
| Construct | Meaning |
|
||||
| --------- | ------- |
|
||||
| `<ppp:ext $lora mappingname/>` | Mapping without additional triggers |
|
||||
| `<ppp:ext $lora mappingname>inline triggers<ppp:/ext>` | Mapping with additional triggers |
|
||||
| Construct | Meaning |
|
||||
|--------------------------------------------------------|-------------------------------------|
|
||||
| `<ppp:ext $lora mappingname/>` | Mapping without 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:
|
||||
|
||||
@@ -480,7 +492,7 @@ See the file in the tests folder as an example.
|
||||
The new format for this command is like this:
|
||||
|
||||
| Construct | Meaning |
|
||||
| --------- | ------- |
|
||||
|---------------------------------------|--------------------------------------------------------------------------------------|
|
||||
| `<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 |
|
||||
|
||||
|
||||
+5
-4
@@ -165,27 +165,28 @@ extranetworktag: "<" /(?!ppp:)\w+:/ encontent ">"
|
||||
choicesoptions_sep: "$$" plain
|
||||
|
||||
// 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
|
||||
choicelabels: ( /"/ IDENTIFIER ( _WHITESPACE? "," _WHITESPACE? IDENTIFIER )* /"/ )
|
||||
| ( /'/ IDENTIFIER ( _WHITESPACE? "," _WHITESPACE? IDENTIFIER )* /'/ )
|
||||
choiceweight: NUMBER
|
||||
choiceif: "if" _WHITESPACE condition
|
||||
choiceelse: "else"
|
||||
choicevalue: content_choice
|
||||
//#endif
|
||||
|
||||
//#if ALLOW_CHOICES
|
||||
// choices construct
|
||||
choices: "{" [ choicesoptions_sampler | ( choicesoptions _WHITESPACE? "$$" ) ] choice ( "|" choice )* "}"
|
||||
choices: "{" [ choicesoptions_sampler | ( choicesoptions "$$" ) ] choice ( "|" choice )* "}"
|
||||
//#endif
|
||||
|
||||
//#if ALLOW_WILDCARDS
|
||||
// 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
|
||||
|
||||
// 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)
|
||||
//#if ALLOW_COMMVARS
|
||||
wildcard_name.2: ( WC_NAME_PLAIN_START | variableuse | commandecho ) ( WC_NAME_PLAIN | variableuse | commandecho )*
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
"""
|
||||
Install dependencies for Prompt Post-Processor extension. For A1111 hosts.
|
||||
"""
|
||||
from pathlib import Path
|
||||
|
||||
requirements_filename = str(Path(__file__).resolve().parent / "requirements.txt")
|
||||
@@ -9,3 +12,5 @@ try:
|
||||
except ImportError:
|
||||
import launch
|
||||
launch.run_pip(f'install -r "{requirements_filename}"', "requirements for Prompt Post-Processor")
|
||||
except Exception as e:
|
||||
pass
|
||||
|
||||
@@ -20,6 +20,9 @@ from ppp_classes import (
|
||||
HostConfig,
|
||||
ModelConfig,
|
||||
ModelDetectConfig,
|
||||
PPPEnvInfo,
|
||||
PPPException,
|
||||
RUN_MODE,
|
||||
VariantConfig,
|
||||
PPPConfig,
|
||||
IFWILDCARDS_CHOICES,
|
||||
@@ -33,7 +36,15 @@ from ppp_variables import VariableRepository, VariableEntry, VariableValue
|
||||
from ppp_logging import DEBUG_LEVEL, log
|
||||
from ppp_tree import TreeProcessor
|
||||
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_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_REMOVE_EXTRANETWORK_TAGS = defopt["cup_remove_extranetwork_tags"]
|
||||
DEFAULT_STRICT_OPERATORS = defopt["strict_operators"]
|
||||
DEFAULT_DO_COMBINATORIAL = defopt["do_combinatorial"]
|
||||
DEFAULT_COMBINATORIAL_SHUFFLE = defopt["combinatorial_shuffle"]
|
||||
DEFAULT_COMBINATORIAL_LIMIT = defopt["combinatorial_limit"]
|
||||
DEFAULT_RUN_MODE = defopt["run_mode"].value
|
||||
DEFAULT_RESULTS_SHUFFLE = defopt["results_shuffle"]
|
||||
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_NEXT_SEED = defopt["next_seed"].value
|
||||
|
||||
WILDCARD_WARNING = '(WARNING TEXT "INVALID WILDCARD" IN BRIGHT RED:1.5)\nBREAK '
|
||||
WILDCARD_STOP = "INVALID WILDCARD! {0}\nBREAK "
|
||||
@@ -82,7 +96,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
def __init__(
|
||||
self,
|
||||
logger: logging.Logger,
|
||||
env_info: dict[str, Any],
|
||||
env_info: PPPEnvInfo,
|
||||
options: PPPStateOptions,
|
||||
grammar_content: Optional[str] = None,
|
||||
interrupt: Optional[Callable] = None,
|
||||
@@ -95,7 +109,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
Args:
|
||||
logger: The logger object.
|
||||
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.
|
||||
grammar_content: Optional. The grammar content to be used for parsing.
|
||||
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):
|
||||
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."""
|
||||
main_folder = Path(__file__).resolve().parent
|
||||
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)
|
||||
self.log(logging.WARNING, errmsg)
|
||||
|
||||
app = env_info.get("app", "")
|
||||
user_config_file = env_info.get("ppp_config", "")
|
||||
app = env_info.app
|
||||
user_config_file = env_info.ppp_config or ""
|
||||
if isinstance(user_config_file, dict):
|
||||
user_cfg, _ = self.__parse_configuration(user_config_file, "forced configuration")
|
||||
else:
|
||||
if user_config_file == "":
|
||||
if app == SUPPORTED_APPS.comfyui.value:
|
||||
if app == SUPPORTED_APPS.comfyui:
|
||||
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()
|
||||
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())
|
||||
|
||||
# Patch for tests (copy comfyui)
|
||||
if app == "tests":
|
||||
if app == SUPPORTED_APPS.tests:
|
||||
if self.config.hosts is None:
|
||||
self.config.hosts = {}
|
||||
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.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:
|
||||
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
|
||||
@@ -394,41 +408,48 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
)
|
||||
self.log(
|
||||
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,
|
||||
)
|
||||
|
||||
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."""
|
||||
prop_base = env_info.get("property_base", None)
|
||||
model_class = env_info.get("model_class", "")
|
||||
model_name = env_info.get("model_filename", "")
|
||||
app = env_info.get("app", "")
|
||||
if not model_class and model_name and app == SUPPORTED_APPS.comfyui.value:
|
||||
model_class = get_model_class_from_filename(model_name)
|
||||
prop_base = env_info.property_base
|
||||
model_class = env_info.model_class
|
||||
model_name = env_info.model_filename
|
||||
app = env_info.app
|
||||
if not model_class and model_name and app == SUPPORTED_APPS.comfyui:
|
||||
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:
|
||||
env_info["model_class"] = model_class
|
||||
env_info.model_class = model_class
|
||||
self.log(
|
||||
logging.DEBUG,
|
||||
f"Detected model class '{model_class}' from filename '{model_name}'",
|
||||
min_level=DEBUG_LEVEL.minimal,
|
||||
)
|
||||
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_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:
|
||||
cls_list = model_detect_for_app.class_ or []
|
||||
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:
|
||||
prop = model_detect_for_app.property
|
||||
attr = getattr(prop_base, prop, None)
|
||||
if isinstance(attr, bool) and attr:
|
||||
env_info["is_" + m] = True
|
||||
env_info.is_flags[m] = True
|
||||
|
||||
def __on_model_info_update(self) -> None:
|
||||
"""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(
|
||||
self,
|
||||
env_info: dict[str, Any],
|
||||
env_info: PPPEnvInfo,
|
||||
options: PPPStateOptions,
|
||||
wildcards_obj: PPPWildcards,
|
||||
extranetwork_mappings_obj: PPPExtraNetworkMappings,
|
||||
@@ -618,14 +639,21 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
)
|
||||
|
||||
@property
|
||||
def envinfo_hash(self) -> str:
|
||||
"""
|
||||
Generates a hash string based on the environment information.
|
||||
|
||||
Returns:
|
||||
str: A hash string representing the environment information.
|
||||
"""
|
||||
return hash(tuple(sorted(self.state.env_info.items())))
|
||||
def envinfo_hash(self) -> int:
|
||||
"""Returns a hash of the environment information for cache-invalidation purposes."""
|
||||
ei = self.state.env_info
|
||||
ppp_config_key = ei.ppp_config if isinstance(ei.ppp_config, (str, type(None))) else id(ei.ppp_config)
|
||||
return hash(
|
||||
(
|
||||
ei.app,
|
||||
ppp_config_key,
|
||||
ei.model_class,
|
||||
ei.model_filename,
|
||||
ei.models_path,
|
||||
id(ei.property_base),
|
||||
tuple(sorted(ei.is_flags.items())),
|
||||
)
|
||||
)
|
||||
|
||||
@property
|
||||
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])
|
||||
|
||||
# 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,
|
||||
# giving a well-defined empty-string fallback without a separate None check.
|
||||
sdchecks.update({"": True})
|
||||
model_name_val = next((k for k, v in sdchecks.items() if v), "")
|
||||
vs.set_system("_model", model_name_val)
|
||||
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("_modelfullname", model_filename)
|
||||
vs.set_system("_sdname", Path(model_filename).name) # deprecated
|
||||
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 = {}
|
||||
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
|
||||
@@ -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_variant_" + x, sdchecks[x] and any(is_models.values()))
|
||||
# special cases
|
||||
vs.set_system("_is_sd", sdchecks.get("sd1", False) or sdchecks.get("sd2", False) or sdchecks.get("sdxl", False) or sdchecks.get("sd3", False))
|
||||
is_ssd = self.state.env_info.get("is_ssd", False)
|
||||
vs.set_system(
|
||||
"_is_sd",
|
||||
sdchecks.get("sd1", False)
|
||||
or sdchecks.get("sd2", False)
|
||||
or sdchecks.get("sdxl", False)
|
||||
or sdchecks.get("sd3", False),
|
||||
)
|
||||
is_ssd = self.state.env_info.is_flags.get("ssd", False)
|
||||
vs.set_system("_is_ssd", 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)
|
||||
@@ -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]}")
|
||||
elif break_processing == "error":
|
||||
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:
|
||||
# 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())
|
||||
for k in var_keys:
|
||||
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)
|
||||
ev = entry.last_echoed_evaluated_value if entry.last_echoed_evaluated_value is not None else entry.value
|
||||
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
|
||||
compound_prompt = prompt + "\n" + negative_prompt
|
||||
found_sequences = re.findall(r"::|\$\$|\$\{|[{}]", compound_prompt)
|
||||
found_sequences = re.findall(r"::|\$\$|\$\{|[{}]|__", compound_prompt)
|
||||
if found_sequences:
|
||||
s = ", ".join(map(lambda x: '"' + x + '"', set(found_sequences)))
|
||||
warnings.append(f"Probably invalid character sequences: {s}.")
|
||||
@@ -1014,8 +1056,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
self,
|
||||
prompt: str,
|
||||
negative_prompt: str,
|
||||
seed: int,
|
||||
starting_seed: int | list[int],
|
||||
jobinfo: Any = None,
|
||||
input_vars: dict[str, Any] | None = None,
|
||||
) -> list[tuple[str, str, dict[str, Any]]]:
|
||||
"""
|
||||
Process the prompt and negative prompt.
|
||||
@@ -1023,7 +1066,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
Args:
|
||||
prompt (str): The 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:
|
||||
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()
|
||||
|
||||
# We update the input state
|
||||
# Truncate the seed to the host's configured bit width (-1 because we only want
|
||||
# positive numbers) so the value stays within the range the host expects
|
||||
# (e.g., 32-bit for SD-WebUI, 64-bit for ComfyUI).
|
||||
self.state.inputs.seed = int(seed & ((1 << (self.state.host_config.seed_bits - 1)) - 1))
|
||||
self.state.inputs.seed = (
|
||||
clamp_host_bits(self.state.host_config.seed_bits, starting_seed)
|
||||
if not isinstance(starting_seed, list)
|
||||
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.neg_prompt = negative_prompt
|
||||
self.state.inputs.jobinfo = jobinfo
|
||||
|
||||
# Input related system variables
|
||||
if input_vars:
|
||||
for k, v in input_vars.items():
|
||||
self.state.variables.set_system("_input_" + k, v)
|
||||
for input_name in self.state.inputs.__dict__.keys():
|
||||
input_value = getattr(self.state.inputs, input_name)
|
||||
var_name = "_input_" + input_name
|
||||
@@ -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_")}
|
||||
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
|
||||
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]]] = []
|
||||
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}:")
|
||||
elif self.state.options.run_mode == RUN_MODE.multiple:
|
||||
self.log(logging.INFO, f"Result {i + 1}:")
|
||||
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)}")
|
||||
if self.state.options.combinatorial_shuffle:
|
||||
rng.shuffle(final_results)
|
||||
self.log(logging.INFO, "Combinations shuffled")
|
||||
if self.state.options.results_shuffle:
|
||||
rng.shuffle(final_results)
|
||||
self.log(logging.INFO, "Results shuffled")
|
||||
return final_results
|
||||
|
||||
def process_prompts_group_start(self):
|
||||
"""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_")}
|
||||
self.log(logging.DEBUG, f"System variables: {filtered_sysvars}")
|
||||
self.log(logging.INFO, f"Combinatorial: {self.state.options.do_combinatorial}")
|
||||
filtered_sysvars = {
|
||||
k: v for k, v in self.state.variables.all_system.items() if not k.startswith(("_input_", "_output_"))
|
||||
}
|
||||
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:
|
||||
"""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"%date%": now.strftime(r"%Y-%m-%d"),
|
||||
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)
|
||||
for token, value in substitutions.items():
|
||||
@@ -1145,11 +1209,16 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
"system_variables": {
|
||||
k: v
|
||||
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": {
|
||||
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},
|
||||
"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,
|
||||
original_prompt: str,
|
||||
original_negative_prompt: str,
|
||||
seed: int = -1,
|
||||
starting_seed: int | list[int] = -1,
|
||||
jobinfo: Any = None,
|
||||
input_vars: dict[str, Any] | None = None,
|
||||
) -> list[tuple[str, str, dict[str, Any]]]:
|
||||
"""
|
||||
Initializes the random number generator and processes the prompt and negative prompt.
|
||||
@@ -1225,23 +1295,29 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
Args:
|
||||
original_prompt (str): The original 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`.
|
||||
input_vars (dict[str, Any] | None): Optional dictionary of input variables to set before processing.
|
||||
|
||||
Returns:
|
||||
list[tuple[str, str, dict[str, Any]]]: A list of tuples containing the processed prompt, negative prompt and all the prompt variables.
|
||||
"""
|
||||
results: list[tuple[str, str, dict[str, Any]]]
|
||||
try:
|
||||
if seed == -1:
|
||||
seed = np.random.randint(0, 2 ** (self.state.host_config.seed_bits - 1), dtype=np.int64)
|
||||
if isinstance(starting_seed, list):
|
||||
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
|
||||
negative_prompt = original_negative_prompt
|
||||
t1 = time.monotonic_ns()
|
||||
if self.state.cyclical_state.last_prompt_pair != (original_prompt, original_negative_prompt):
|
||||
self.state.cyclical_state.reset()
|
||||
self.state.cyclical_state.last_prompt_pair = (original_prompt, original_negative_prompt)
|
||||
results = self.__processprompts(prompt, negative_prompt, seed, jobinfo)
|
||||
results = self.__processprompts(prompt, negative_prompt, starting_seed, jobinfo, input_vars or {})
|
||||
t2 = time.monotonic_ns()
|
||||
self.log(logging.INFO, f"Process prompt pair time: {(t2 - t1) / 1_000_000_000:.3f} seconds")
|
||||
# self.log(logging.DEBUG,f"Wildcards memory usage: {self.state.wildcards_obj.__sizeof__()}")
|
||||
|
||||
+55
-6
@@ -47,6 +47,24 @@ class ONWARNING_CHOICES(Enum):
|
||||
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 -------------------
|
||||
|
||||
AttentionOption = Literal["ok", "parentheses", "disable", "remove", "error"]
|
||||
@@ -209,10 +227,13 @@ class PPPStateOptions:
|
||||
cup_merge_attention: bool = True
|
||||
cup_remove_extranetwork_tags: bool = False
|
||||
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
|
||||
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):
|
||||
if not self.cup_do_cleanup:
|
||||
@@ -235,7 +256,7 @@ class PPPStateOptions:
|
||||
class PPPStateInputs:
|
||||
"""Structured inputs for a single prompt processing call."""
|
||||
|
||||
seed: int = -1
|
||||
seed: int | list[int] = -1
|
||||
pos_prompt: str = ""
|
||||
neg_prompt: str = ""
|
||||
jobinfo: Any = None
|
||||
@@ -272,12 +293,30 @@ class CyclicalSamplerState:
|
||||
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)
|
||||
class PPPState:
|
||||
"""State object passed to various PPP components during prompt processing."""
|
||||
|
||||
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)
|
||||
options: PPPStateOptions = field(default_factory=PPPStateOptions)
|
||||
inputs: PPPStateInputs = field(default_factory=PPPStateInputs)
|
||||
@@ -288,7 +327,17 @@ class PPPState:
|
||||
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.
|
||||
This exception can be raised to stop the processing of prompts.
|
||||
|
||||
+160
-59
@@ -1,17 +1,28 @@
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit("This script must be run from ComfyUI")
|
||||
|
||||
# pylint: disable=wrong-import-position,wrong-import-order
|
||||
from datetime import datetime
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import folder_paths # type: ignore
|
||||
import nodes # type: ignore
|
||||
import folder_paths # type: ignore # pylint: disable=import-error
|
||||
import nodes # type: ignore # pylint: disable=import-error
|
||||
|
||||
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_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory, log
|
||||
from ppp_utils import escape_single_quotes
|
||||
@@ -190,29 +201,20 @@ class PromptPostProcessorComfyUINode:
|
||||
"label_off": "No",
|
||||
},
|
||||
),
|
||||
"do_combinatorial": (
|
||||
"BOOLEAN",
|
||||
"results_file": (
|
||||
"STRING",
|
||||
{
|
||||
"default": PromptPostProcessor.DEFAULT_DO_COMBINATORIAL,
|
||||
"tooltip": "Enable combinatorial mode",
|
||||
"label_on": "Yes",
|
||||
"label_off": "No",
|
||||
"default": PromptPostProcessor.DEFAULT_RESULTS_FILE,
|
||||
"tooltip": r"Filename to save processing results. Supports %datetime%, %date%, %time%, %host% tokens. Empty = disabled.",
|
||||
"dynamicPrompts": False,
|
||||
},
|
||||
),
|
||||
"combinatorial_shuffle": (
|
||||
"BOOLEAN",
|
||||
"run_mode": (
|
||||
"COMBO",
|
||||
{
|
||||
"default": PromptPostProcessor.DEFAULT_COMBINATORIAL_SHUFFLE,
|
||||
"tooltip": "Shuffle the combinatorial results",
|
||||
"label_on": "Yes",
|
||||
"label_off": "No",
|
||||
},
|
||||
),
|
||||
"combinatorial_limit": (
|
||||
"INT",
|
||||
{
|
||||
"default": PromptPostProcessor.DEFAULT_COMBINATORIAL_LIMIT,
|
||||
"tooltip": "Limit for combinatorial mode",
|
||||
"options": [e.value for e in RUN_MODE],
|
||||
"default": PromptPostProcessor.DEFAULT_RUN_MODE,
|
||||
"tooltip": "Run mode",
|
||||
},
|
||||
),
|
||||
"wc_options": (
|
||||
@@ -243,12 +245,11 @@ class PromptPostProcessorComfyUINode:
|
||||
"tooltip": "ExtraNetworks mapping options",
|
||||
},
|
||||
),
|
||||
"results_file": (
|
||||
"STRING",
|
||||
"rm_options": (
|
||||
"PPP_OPTIONS_RM",
|
||||
{
|
||||
"default": PromptPostProcessor.DEFAULT_RESULTS_FILE,
|
||||
"tooltip": r"Filename to save processing results. Supports %datetime%, %date%, %time%, %host% tokens. Empty = disabled.",
|
||||
"dynamicPrompts": False,
|
||||
"default": None,
|
||||
"tooltip": "Run mode options",
|
||||
},
|
||||
),
|
||||
},
|
||||
@@ -283,9 +284,9 @@ class PromptPostProcessorComfyUINode:
|
||||
"variables",
|
||||
)
|
||||
OUTPUT_TOOLTIPS = (
|
||||
"Processed positive prompt (list of prompts if combinatorial mode is enabled)",
|
||||
"Processed negative prompt (list of prompts if combinatorial mode is enabled)",
|
||||
"Output variables (list of dictionaries 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/multiple mode is enabled)",
|
||||
"Output variables (list of dictionaries if combinatorial/multiple mode is enabled)",
|
||||
)
|
||||
|
||||
FUNCTION = "process"
|
||||
@@ -304,27 +305,34 @@ class PromptPostProcessorComfyUINode:
|
||||
seed,
|
||||
debug_level,
|
||||
on_warnings,
|
||||
strict_operators,
|
||||
process_wildcards,
|
||||
do_cleanup,
|
||||
cleanup_variables,
|
||||
do_combinatorial,
|
||||
combinatorial_shuffle,
|
||||
combinatorial_limit,
|
||||
results_file,
|
||||
run_mode,
|
||||
model=None,
|
||||
wc_options=None,
|
||||
stn_options=None,
|
||||
cup_options=None,
|
||||
en_options=None,
|
||||
strict_operators=None,
|
||||
results_file=None,
|
||||
rm_options=None,
|
||||
):
|
||||
modelclass = (
|
||||
model.model.model_config.__class__.__name__ if model is not None and not isinstance(model, str) else model
|
||||
) or ""
|
||||
if modelname == "(none)":
|
||||
modelname = ""
|
||||
if modelclass == "":
|
||||
modelclass = get_model_class_from_filename(modelname)
|
||||
if modelclass == "" and 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:
|
||||
log(
|
||||
self.logger,
|
||||
@@ -337,7 +345,7 @@ class PromptPostProcessorComfyUINode:
|
||||
self.logger,
|
||||
DEBUG_LEVEL.minimal,
|
||||
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 == "":
|
||||
log(
|
||||
@@ -347,24 +355,26 @@ class PromptPostProcessorComfyUINode:
|
||||
"Modelname was not provided. System model and variant variables will not be properly set.",
|
||||
)
|
||||
# model class values in ComfyUI\comfy\supported_models.py
|
||||
env_info = {
|
||||
"app": SUPPORTED_APPS.comfyui.value,
|
||||
"models_path": folder_paths.models_dir,
|
||||
"model_filename": modelname or "", # path is relative to checkpoints folder
|
||||
"model_class": modelclass,
|
||||
"property_base": None,
|
||||
}
|
||||
env_info = PPPEnvInfo(
|
||||
app=SUPPORTED_APPS.comfyui,
|
||||
models_path=folder_paths.models_dir,
|
||||
model_filename=modelname or "", # path is relative to checkpoints folder
|
||||
model_class=modelclass,
|
||||
property_base=None,
|
||||
)
|
||||
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 "")
|
||||
|
||||
options = PPPStateOptions(
|
||||
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 if strict_operators is not None else PromptPostProcessor.DEFAULT_STRICT_OPERATORS
|
||||
),
|
||||
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=(
|
||||
wc_options["wc_choice_separator"] if wc_options else PromptPostProcessor.DEFAULT_CHOICE_SEPARATOR
|
||||
),
|
||||
@@ -415,10 +425,19 @@ class PromptPostProcessorComfyUINode:
|
||||
if cup_options
|
||||
else PromptPostProcessor.DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS
|
||||
),
|
||||
do_combinatorial=do_combinatorial,
|
||||
combinatorial_shuffle=combinatorial_shuffle,
|
||||
combinatorial_limit=combinatorial_limit,
|
||||
results_file=results_file or "",
|
||||
run_mode=RUN_MODE(run_mode if run_mode else PromptPostProcessor.DEFAULT_RUN_MODE),
|
||||
results_file=results_file,
|
||||
results_shuffle=(
|
||||
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(
|
||||
options.debug_level,
|
||||
@@ -455,13 +474,94 @@ class PromptPostProcessorComfyUINode:
|
||||
jobinfo={"job_timestamp": datetime.now().isoformat()},
|
||||
)
|
||||
self.ppp.process_prompts_group_end()
|
||||
|
||||
return tuple(zip(*results)) # unzip the list of tuples into tuple of lists
|
||||
|
||||
def interrupt(self):
|
||||
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:
|
||||
"""
|
||||
Node for wildcard options.
|
||||
@@ -907,13 +1007,13 @@ class PromptPostProcessorWildcardConcatComfyUINode:
|
||||
lf = PromptPostProcessorLogFactory()
|
||||
cls._ppp = PromptPostProcessor(
|
||||
lf.log,
|
||||
{
|
||||
"app": SUPPORTED_APPS.comfyui.value,
|
||||
"models_path": folder_paths.models_dir,
|
||||
"model_filename": "",
|
||||
"model_class": "",
|
||||
"property_base": None,
|
||||
},
|
||||
PPPEnvInfo(
|
||||
app=SUPPORTED_APPS.comfyui,
|
||||
models_path=folder_paths.models_dir,
|
||||
model_filename="",
|
||||
model_class="",
|
||||
property_base=None,
|
||||
),
|
||||
PPPStateOptions(debug_level=DEBUG_LEVEL.minimal),
|
||||
wildcards_obj=PPPWildcards(lf.log),
|
||||
)
|
||||
@@ -1031,6 +1131,7 @@ class PromptPostProcessorWildcardConcatComfyUINode:
|
||||
|
||||
|
||||
try:
|
||||
# pylint: disable=import-error
|
||||
from server import PromptServer # type: ignore
|
||||
from aiohttp import web as _aiohttp_web # type: ignore
|
||||
|
||||
|
||||
+85
-20
@@ -1,5 +1,6 @@
|
||||
import ast
|
||||
import csv
|
||||
from enum import Enum
|
||||
from functools import reduce
|
||||
import logging
|
||||
from pathlib import Path
|
||||
@@ -10,7 +11,7 @@ import lark
|
||||
from ruamel.yaml import YAML as _YAML
|
||||
|
||||
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
|
||||
|
||||
|
||||
@@ -82,13 +83,19 @@ def parse_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 "
|
||||
if state.options.on_warning == ONWARNING_CHOICES.stop:
|
||||
raise PPPInterrupt(
|
||||
message,
|
||||
INVALID_CONTENT_STOP.format(message) if not is_negative else "",
|
||||
INVALID_CONTENT_STOP.format(message) if is_negative else "",
|
||||
INVALID_CONTENT_STOP.format(message) if where == WARN_STOP_WHERE.positive else "",
|
||||
INVALID_CONTENT_STOP.format(message) if where == WARN_STOP_WHERE.negative else "",
|
||||
) from e
|
||||
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)
|
||||
|
||||
|
||||
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:
|
||||
import folder_paths # type: ignore
|
||||
import comfy.utils # type: ignore
|
||||
import comfy.model_detection as model_detection # type: ignore
|
||||
except ImportError:
|
||||
return ""
|
||||
except ImportError as e:
|
||||
raise PPPException(f"Error detecting class from '{filename}': {e}") from e
|
||||
import json
|
||||
|
||||
if not filename:
|
||||
return ""
|
||||
full_path = (
|
||||
folder_paths.get_full_path("diffusion_models", filename)
|
||||
or folder_paths.get_full_path("checkpoints", filename)
|
||||
or folder_paths.get_full_path("unet", filename)
|
||||
)
|
||||
if not full_path or not full_path.lower().endswith((".safetensors", ".sft")):
|
||||
return ""
|
||||
return None
|
||||
|
||||
path_keys = ["diffusion_models", "checkpoints", "unet"]
|
||||
full_path: Path | None = None
|
||||
if filename.is_absolute():
|
||||
full_path = filename
|
||||
else:
|
||||
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:
|
||||
header_bytes = comfy.utils.safetensors_header(full_path)
|
||||
if header_bytes is None:
|
||||
return ""
|
||||
raise PPPException(f"Error detecting class from '{full_path}': no header")
|
||||
header = json.loads(header_bytes)
|
||||
|
||||
# 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}
|
||||
|
||||
prefix = model_detection.unet_prefix_from_state_dict(mock_sd)
|
||||
config = model_detection.model_config_from_unet(mock_sd, prefix)
|
||||
return config.__class__.__name__ if config else ""
|
||||
except Exception: # pylint: disable=broad-except
|
||||
config = model_detection.model_config_from_unet(mock_sd, prefix, True)
|
||||
if not config:
|
||||
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 ""
|
||||
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:
|
||||
@@ -288,7 +339,7 @@ def convert_sdnext_styles_to_wildcard(inp: Path, out: Path):
|
||||
files = inp.glob("*.json") if inp.is_dir() else [inp]
|
||||
for file in files:
|
||||
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):
|
||||
continue
|
||||
wildcards[file] = {}
|
||||
@@ -309,3 +360,17 @@ def convert_sdnext_styles_to_wildcard(inp: Path, out: Path):
|
||||
f.write(f"# Converted from {name}\n")
|
||||
yaml_writer = _YAML()
|
||||
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))
|
||||
|
||||
@@ -82,7 +82,7 @@ hosts:
|
||||
alternation: error
|
||||
and: comma
|
||||
break: comma
|
||||
seed_bits: 64
|
||||
seed_bits: 53 # and not 64, because of frontend JavaScript limits
|
||||
|
||||
# Supported base models, variants, and options
|
||||
# Check supported models for each host in:
|
||||
@@ -101,7 +101,7 @@ models:
|
||||
a1111: { property: "is_sd1" }
|
||||
forge: { property: "is_sd1", class: ["SD15", "SD15_instructpix2pix"] }
|
||||
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
|
||||
comfyui: { class: ["SD15", "SD15_instructpix2pix"] }
|
||||
sd2: # Stable Diffusion 2
|
||||
@@ -109,7 +109,7 @@ models:
|
||||
a1111: { property: "is_sd2" }
|
||||
forge: { property: "is_sd2", class: ["SD20", "SD21UnclipL", "SD21UnclipH"] }
|
||||
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
|
||||
comfyui: { class: ["SD20", "SD21UnclipL", "SD21UnclipH", "LotusD"] }
|
||||
ssd: # Segmind Stable Diffusion 1B
|
||||
@@ -117,7 +117,7 @@ models:
|
||||
a1111: { property: "is_ssd" }
|
||||
forge: { class: ["SSD1B"] }
|
||||
forgeneo: null
|
||||
reforge: { property: "is_ssd" }
|
||||
reforge: { property: "is_ssd", class: ["SSD1B"] }
|
||||
sdnext: null
|
||||
comfyui: { class: ["SSD1B"]}
|
||||
sdxl: # Stable Diffusion XL
|
||||
@@ -125,7 +125,7 @@ models:
|
||||
a1111: { property: "is_sdxl" }
|
||||
forge: { property: "is_sdxl", class: ["SDXL", "SDXLRefiner", "SDXL_instructpix2pix", "Segmind_Vega", "KOALA_700M", "KOALA_1B"] }
|
||||
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"] }
|
||||
comfyui: { class: ["SDXL", "SDXLRefiner", "SDXL_instructpix2pix", "Segmind_Vega", "KOALA_700M", "KOALA_1B"] }
|
||||
variants:
|
||||
@@ -142,7 +142,7 @@ models:
|
||||
a1111: { property: "is_sd3" }
|
||||
forge: { property: "is_sd3", class: ["SD3"] }
|
||||
forgeneo: null
|
||||
reforge: { property: "is_sd3" }
|
||||
reforge: { property: "is_sd3", class: ["SD3"] }
|
||||
sdnext: { class: ["StableDiffusion3Pipeline"] }
|
||||
comfyui: { class: ["SD3"] }
|
||||
flux: # Flux 1
|
||||
@@ -190,7 +190,7 @@ models:
|
||||
a1111: null
|
||||
forge: null
|
||||
forgeneo: null
|
||||
reforge: null
|
||||
reforge: { class: ["LTXV"] }
|
||||
sdnext: null
|
||||
comfyui: { class: ["LTXV", "LTXAV"] }
|
||||
cosmos: # Cosmos
|
||||
@@ -248,7 +248,7 @@ models:
|
||||
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"] }
|
||||
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
|
||||
detect:
|
||||
a1111: null
|
||||
@@ -337,3 +337,59 @@ models:
|
||||
reforge: null
|
||||
sdnext: null
|
||||
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
@@ -218,7 +218,7 @@ class PPPExtraNetworkMappings:
|
||||
enmappings_input = enmappings_input.strip()
|
||||
if enmappings_input != "":
|
||||
try:
|
||||
content = _YAML(typ='safe').load(enmappings_input)
|
||||
content = _YAML(typ="safe").load(enmappings_input)
|
||||
except _YAMLError as e:
|
||||
log(
|
||||
self.__logger,
|
||||
@@ -298,7 +298,7 @@ class PPPExtraNetworkMappings:
|
||||
try:
|
||||
try:
|
||||
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
|
||||
log(
|
||||
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...",
|
||||
)
|
||||
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)
|
||||
except Exception as e: # pylint: disable=broad-except
|
||||
log(
|
||||
|
||||
+222
-109
@@ -11,11 +11,11 @@ from typing import Callable, Optional
|
||||
import lark
|
||||
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_logging import DEBUG_LEVEL, log
|
||||
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_wildcards import PPPWildcard
|
||||
|
||||
@@ -62,6 +62,9 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
self.__on_model_info_update = on_model_info_update
|
||||
self.__debug_level = state.options.debug_level
|
||||
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.__negtags: list[TreeProcessor.NegTag] = [] # type: ignore
|
||||
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.__insertion_at: list[tuple[int, int]] = [None for _ in range(10)]
|
||||
self.__detectedWildcards: list[tuple[str, bool]] = []
|
||||
self.__current_run_index = 0
|
||||
self.__result = ""
|
||||
self.__comb_forced_path: list[int] = []
|
||||
self.__comb_trace: list[int] = []
|
||||
self.__cycl_forced_path: list[int] = []
|
||||
self.__cycl_trace: list[int] = []
|
||||
self.__forced_path: list[int] = []
|
||||
self.__trace: list[int] = []
|
||||
self.__rand_decisions: dict[int, int] = {}
|
||||
|
||||
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)
|
||||
|
||||
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):
|
||||
"""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.__negtags = []
|
||||
self.__already_processed = []
|
||||
@@ -98,6 +106,35 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
if self.state.extranetwork_mappings_obj is not None:
|
||||
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(
|
||||
self,
|
||||
parsed: lark.Tree,
|
||||
@@ -112,54 +149,89 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
Returns:
|
||||
list[tuple[str, list[tuple[str,bool]], dict[str, VariableEntry]]]: A list of
|
||||
(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.
|
||||
"""
|
||||
self.log(logging.INFO, "Processing prompt...")
|
||||
|
||||
self.__detectedWildcards = []
|
||||
self.__is_negative = False
|
||||
self.__result = ""
|
||||
results: list[tuple[str, list[tuple[str, bool]], tuple]] = []
|
||||
max_results = (
|
||||
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:
|
||||
self.__cycl_forced_path = list(self.state.cyclical_state.current_path)
|
||||
self.__cycl_trace = []
|
||||
self.visit(parsed)
|
||||
self.__finalize_variables()
|
||||
if self.__cycl_trace:
|
||||
self.state.cyclical_state.last_trace = self.__cycl_trace[:]
|
||||
self.state.cyclical_state.advance()
|
||||
return [(self.__result, self.__detectedWildcards, self.state.variables.backup_user())]
|
||||
initial_vars = self.state.variables.backup_user()
|
||||
|
||||
if self.state.options.run_mode != RUN_MODE.combinatorial:
|
||||
initial_path = list(self.state.cyclical_state.current_path)
|
||||
warned_cycle = False
|
||||
for step in range(max_results):
|
||||
self.log(logging.DEBUG, f"Using seed {self.__current_input_seed}")
|
||||
self.__reset_run_state()
|
||||
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.
|
||||
# __comb_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
|
||||
# __forced_path drives which option is selected at each decision point;
|
||||
# __trace records how many options were available at each point so the DFS can
|
||||
# correctly enumerate unexplored branches after each run.
|
||||
initial_vars = self.state.variables.backup_user()
|
||||
results: list[tuple[str, list[tuple[str, bool]], tuple]] = []
|
||||
limit = self.state.options.combinatorial_limit
|
||||
# __rand_decisions caches random (~) choices so they stay consistent across all runs.
|
||||
self.__rand_decisions = {}
|
||||
|
||||
def _run(forced_path: tuple[int, ...]) -> tuple[int, ...]:
|
||||
self.log(logging.DEBUG, f"Running combinatorial path: {forced_path}")
|
||||
self.__comb_forced_path = list(forced_path)
|
||||
self.__comb_trace = []
|
||||
self.log(logging.DEBUG, f"Using seed {self.__current_input_seed}")
|
||||
self.__forced_path = list(forced_path)
|
||||
self.__trace = []
|
||||
self.__reset_run_state()
|
||||
self.state.variables.restore_user(initial_vars)
|
||||
self.visit(parsed)
|
||||
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:
|
||||
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"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
|
||||
|
||||
def _dfs(forced_path: tuple[int, ...]):
|
||||
"""Recursively explore combinatorial branches via Depth First Search (DFS)."""
|
||||
nonlocal limit_reached
|
||||
if 0 < limit <= len(results):
|
||||
if 0 < max_results <= len(results):
|
||||
limit_reached = True
|
||||
return
|
||||
trace = _run(forced_path)
|
||||
@@ -168,12 +240,12 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
# Iterating in reverse means later (deeper) decision points vary fastest,
|
||||
# so the output order is depth-first rather than breadth-first.
|
||||
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
|
||||
return
|
||||
num_options = trace[i]
|
||||
for opt in range(1, num_options):
|
||||
if 0 < limit <= len(results):
|
||||
if 0 < max_results <= len(results):
|
||||
limit_reached = True
|
||||
return
|
||||
# Pad with zeros for intermediate decisions so they keep the default.
|
||||
@@ -182,7 +254,9 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
|
||||
_dfs(())
|
||||
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
|
||||
|
||||
def __finalize_variables(self):
|
||||
@@ -198,6 +272,8 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
name, specifier = self.__separate_arrayref(k)
|
||||
value = self.get_final_scalar_variable(name, specifier)
|
||||
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(
|
||||
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."
|
||||
)
|
||||
return
|
||||
app = self.state.env_info.get("app", "")
|
||||
if app not in (SUPPORTED_APPS.comfyui.value, SUPPORTED_APPS.tests.value):
|
||||
app = self.state.env_info.app
|
||||
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.")
|
||||
return
|
||||
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":
|
||||
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:
|
||||
self.__on_model_info_update()
|
||||
info = variable_name + " = " + f"'{escape_single_quotes(evaluated)}'"
|
||||
@@ -1641,24 +1717,38 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
else_mapping = v
|
||||
num_mappings = len(found_mappings)
|
||||
if num_mappings > 0:
|
||||
if self.state.options.do_combinatorial:
|
||||
decision_idx = len(self.__comb_trace)
|
||||
self.__comb_trace.append(num_mappings)
|
||||
if num_mappings == 1:
|
||||
found = found_mappings[0]
|
||||
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 = (
|
||||
min(self.__comb_forced_path[decision_idx], num_mappings - 1)
|
||||
if decision_idx < len(self.__comb_forced_path)
|
||||
min(self.__forced_path[decision_idx], num_mappings - 1)
|
||||
if decision_idx < len(self.__forced_path)
|
||||
else 0
|
||||
)
|
||||
found = found_mappings[chosen_idx]
|
||||
elif num_mappings == 1:
|
||||
found = found_mappings[0]
|
||||
else:
|
||||
found = found_mappings[
|
||||
self.__rng.choice(
|
||||
num_mappings,
|
||||
p=[v.weight or 1 for v in found_mappings],
|
||||
)
|
||||
]
|
||||
elif self.state.options.default_sampler == DEFAULT_SAMPLER.cyclical:
|
||||
decision_idx = len(self.__trace)
|
||||
self.__trace.append(num_mappings)
|
||||
chosen_idx = (
|
||||
self.__forced_path[decision_idx] % num_mappings
|
||||
if decision_idx < len(self.__forced_path)
|
||||
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:
|
||||
found = else_mapping
|
||||
# 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)
|
||||
if options is None:
|
||||
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)
|
||||
optional: bool = options.get("optional", False)
|
||||
if "count" in options:
|
||||
@@ -1891,8 +1986,12 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
included_choices = 0
|
||||
excluded_choices = 0
|
||||
excluded_weights_sum = 0
|
||||
else_choice = None
|
||||
for i, c in enumerate(expanded_choice_values):
|
||||
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))
|
||||
condition = c.get("if", None)
|
||||
if weight > 0 and (condition is None or self.__eval_condition(condition)):
|
||||
@@ -1903,10 +2002,15 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
weights.append(-1)
|
||||
excluded_choices += 1
|
||||
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
|
||||
weights = [weight + excluded_weights_sum / included_choices for weight in weights if weight >= 0]
|
||||
weights = np.array(weights)
|
||||
weights /= weights.sum() # normalize weights
|
||||
selected_choices: list[dict] = []
|
||||
if available_choices:
|
||||
if from_value < 0:
|
||||
from_value = 1
|
||||
@@ -1916,8 +2020,7 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
to_value = 1
|
||||
elif (to_value > len(available_choices) and not repeating) or from_value > to_value:
|
||||
to_value = len(available_choices)
|
||||
comb_chosen_selection: Optional[list[dict]] = None
|
||||
if self.state.options.do_combinatorial or sampler == "@":
|
||||
if self.state.options.run_mode == RUN_MODE.combinatorial or sampler == "@":
|
||||
# Enumerate every distinct selection of choices, accounting for count range and repetition.
|
||||
all_selections: list[tuple] = []
|
||||
# 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))
|
||||
else:
|
||||
all_selections.extend(permutations(available_choices, k))
|
||||
num_selections = len(all_selections)
|
||||
if self.state.options.do_combinatorial:
|
||||
decision_idx = len(self.__comb_trace)
|
||||
self.__comb_trace.append(num_selections)
|
||||
decision_idx = len(self.__trace)
|
||||
if self.state.options.run_mode == RUN_MODE.combinatorial:
|
||||
if specified_sampler == "~":
|
||||
# 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 = (
|
||||
min(self.__comb_forced_path[decision_idx], num_selections - 1)
|
||||
if decision_idx < len(self.__comb_forced_path)
|
||||
min(self.__forced_path[decision_idx], num_selections - 1)
|
||||
if decision_idx < len(self.__forced_path)
|
||||
else 0
|
||||
)
|
||||
else: # sampler == "@"
|
||||
cycl_decision_idx = len(self.__cycl_trace)
|
||||
self.__cycl_trace.append(num_selections)
|
||||
num_selections = len(all_selections)
|
||||
self.__trace.append(num_selections)
|
||||
chosen_idx = (
|
||||
self.__cycl_forced_path[cycl_decision_idx] % num_selections
|
||||
if cycl_decision_idx < len(self.__cycl_forced_path)
|
||||
self.__forced_path[decision_idx] % num_selections
|
||||
if decision_idx < len(self.__forced_path)
|
||||
else 0
|
||||
)
|
||||
comb_chosen_selection = list(all_selections[chosen_idx])
|
||||
num_choices = len(comb_chosen_selection)
|
||||
selected_choices = list(all_selections[chosen_idx])
|
||||
else:
|
||||
num_choices = (
|
||||
chosen_num_choices = (
|
||||
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:
|
||||
num_choices = 0
|
||||
if not optional and from_value > 0:
|
||||
self.warn_or_stop(f"Not enough choices found for {msg_where}!")
|
||||
if num_choices < 2:
|
||||
repeating = False
|
||||
if self.state.options.keep_choices_order:
|
||||
selected_choices = sorted(selected_choices, key=lambda x: x["choice_index"])
|
||||
num_choices = len(selected_choices)
|
||||
self.log(
|
||||
logging.DEBUG,
|
||||
f"Selecting {'optional ' if optional else ''}{'repeating ' if repeating else ''}{num_choices} choice"
|
||||
+ ("s" if num_choices != 1 else "")
|
||||
+ (f" and separating with '{escape_single_quotes(separator)}'" if num_choices > 1 else ""),
|
||||
)
|
||||
if num_choices > 0:
|
||||
if comb_chosen_selection is not None:
|
||||
selected_choices: list[dict] = comb_chosen_selection
|
||||
selected_choices_text = []
|
||||
for i, c in enumerate(selected_choices):
|
||||
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:
|
||||
selected_choices: list[dict] = (
|
||||
list(self.__rng.choice(available_choices, size=num_choices, p=weights, replace=repeating))
|
||||
if available_choices
|
||||
else []
|
||||
)
|
||||
if self.state.options.keep_choices_order:
|
||||
selected_choices = sorted(selected_choices, key=lambda x: x["choice_index"])
|
||||
selected_choices_text = []
|
||||
for i, c in enumerate(selected_choices):
|
||||
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:
|
||||
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 = []
|
||||
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]
|
||||
container = options.get("container", None)
|
||||
if container is None:
|
||||
separator = options.get("separator", self.state.options.choice_separator)
|
||||
@@ -2030,11 +2136,12 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
),
|
||||
],
|
||||
)
|
||||
self.log(
|
||||
logging.DEBUG,
|
||||
"Unseen wildcards: "
|
||||
+ ", ".join([f"'{escape_single_quotes(x)}'" for x in self.__seen_wildcards[seen_wildcards_len:]]),
|
||||
)
|
||||
if self.__seen_wildcards[seen_wildcards_len:]:
|
||||
self.log(
|
||||
logging.DEBUG,
|
||||
"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]
|
||||
return container, results
|
||||
|
||||
@@ -2114,7 +2221,12 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
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["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]
|
||||
return choice_dict
|
||||
|
||||
@@ -2242,7 +2354,7 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
parse_prompt(
|
||||
self.state,
|
||||
"as wildcard options",
|
||||
wildcard.unprocessed_choices[0][:-2].strip(),
|
||||
wildcard.unprocessed_choices[0],
|
||||
self.state.parsers["wcdefoptions"],
|
||||
True,
|
||||
),
|
||||
@@ -2371,7 +2483,8 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
self.__result += wc
|
||||
if self.__debug_level == DEBUG_LEVEL.full:
|
||||
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]
|
||||
t2 = time.monotonic_ns()
|
||||
self.__debug_end("wildcard", start_result, t2 - t1, f"'{escape_single_quotes(wc)}'")
|
||||
|
||||
@@ -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 ----
|
||||
|
||||
def get(self, name: str, default: Any = None) -> Any:
|
||||
|
||||
+46
-9
@@ -1,5 +1,6 @@
|
||||
import fnmatch
|
||||
from pathlib import Path
|
||||
import re
|
||||
from typing import Any, Optional
|
||||
import logging
|
||||
from ruamel.yaml import YAML as _YAML
|
||||
@@ -114,23 +115,33 @@ class PPPWildcards:
|
||||
keys = sorted(fnmatch.filter(self.wildcards.keys(), key))
|
||||
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.
|
||||
|
||||
Args:
|
||||
dictionary (dict): The dictionary to check.
|
||||
prefix (str): The prefix for the current key.
|
||||
file_str (str): The file string for logging purposes.
|
||||
|
||||
Returns:
|
||||
list: A list of all leaf wildcards in the dictionary.
|
||||
"""
|
||||
wc = []
|
||||
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):
|
||||
wc.extend(self.__get_wc_in_dict(obj, prefix + str(key) + "/"))
|
||||
wc.extend(self.__get_wc_in_dict(obj, prefix + strkey + "/", file_str))
|
||||
else:
|
||||
wc.append((prefix + str(key), obj))
|
||||
wc.append((prefix + strkey, obj))
|
||||
return wc
|
||||
|
||||
def __remove_wildcards_from_path(self, full_path: Path, debug=True):
|
||||
@@ -218,7 +229,7 @@ class PPPWildcards:
|
||||
wildcards_input = wildcards_input.strip()
|
||||
if wildcards_input != "":
|
||||
try:
|
||||
content = _YAML(typ='safe').load(wildcards_input)
|
||||
content = _YAML(typ="safe").load(wildcards_input)
|
||||
except _YAMLError as e:
|
||||
log(self.__logger, self.__debug_level, logging.WARNING, f"Invalid format for input wildcards: {e}")
|
||||
return
|
||||
@@ -268,7 +279,7 @@ class PPPWildcards:
|
||||
Returns:
|
||||
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:
|
||||
"""
|
||||
@@ -374,6 +385,21 @@ class PPPWildcards:
|
||||
value = f"{options}::{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]):
|
||||
"""
|
||||
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"
|
||||
|
||||
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):
|
||||
key_parts.pop()
|
||||
keys = self.__get_wc_in_dict(content)
|
||||
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, "", file_str)
|
||||
for key, obj in keys:
|
||||
tmp_key_parts = key_parts.copy()
|
||||
tmp_key_parts.extend(key.split("/"))
|
||||
@@ -472,7 +509,7 @@ class PPPWildcards:
|
||||
external_key_parts = list(full_path.with_suffix("").relative_to(base).parts)
|
||||
try:
|
||||
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
|
||||
log(
|
||||
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...",
|
||||
)
|
||||
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)
|
||||
|
||||
def __get_wildcards_in_text_file(self, full_path: Path, base: Path):
|
||||
|
||||
+8
-7
@@ -1,18 +1,19 @@
|
||||
[project]
|
||||
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."
|
||||
version = "3.2.1"
|
||||
version = "3.3.0"
|
||||
license = { file = "LICENSE.txt" }
|
||||
dependencies = ["lark", "numpy", "ruamel.yaml", "pydantic"]
|
||||
dependencies = ["lark==1.*", "numpy==2.*", "ruamel.yaml==0.*", "pydantic==2.*"]
|
||||
requires-python = ">=3.10"
|
||||
|
||||
[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/
|
||||
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"
|
||||
|
||||
[tool.comfy]
|
||||
publisher_id = "acorderob"
|
||||
display_name = "ACB Prompt PostProcessor"
|
||||
icon = "https://raw.githubusercontent.com/acorderob/sd-webui-prompt-postprocessor/main/images/prompt-postprocessor-icon.png"
|
||||
PublisherId = "acorderob"
|
||||
DisplayName = "ACB Prompt PostProcessor"
|
||||
Icon = "https://raw.githubusercontent.com/acorderob/sd-webui-prompt-postprocessor/main/images/prompt-postprocessor-icon.png"
|
||||
|
||||
+4
-4
@@ -1,4 +1,4 @@
|
||||
lark
|
||||
numpy
|
||||
ruamel.yaml
|
||||
pydantic
|
||||
lark==1.*
|
||||
numpy==2.*
|
||||
ruamel.yaml==0.*
|
||||
pydantic==2.*
|
||||
+214
-122
@@ -1,6 +1,7 @@
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit("This script must be run from a Stable Diffusion WebUI")
|
||||
|
||||
# pylint: disable=wrong-import-position,wrong-import-order
|
||||
import logging
|
||||
import sys
|
||||
import os
|
||||
@@ -10,14 +11,24 @@ import numpy as np
|
||||
|
||||
sys.path.append(str(Path(__file__).parent)) # base path for the extension
|
||||
|
||||
from modules import scripts, shared, script_callbacks # type: ignore
|
||||
from modules.processing import StableDiffusionProcessing # type: ignore
|
||||
from modules.shared import opts # type: ignore
|
||||
from modules.paths import models_path # type: ignore
|
||||
import gradio as gr # type: ignore
|
||||
from modules import scripts, shared, script_callbacks # type: ignore # pylint: disable=import-error
|
||||
from modules.processing import StableDiffusionProcessing # type: ignore # pylint: disable=import-error
|
||||
from modules.shared import opts # type: ignore # pylint: disable=import-error
|
||||
from modules.paths import models_path # type: ignore # pylint: disable=import-error
|
||||
import gradio as gr # type: ignore # pylint: disable=import-error
|
||||
|
||||
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_cache import PPPLRUCache
|
||||
from ppp_wildcards import PPPWildcards
|
||||
@@ -74,11 +85,13 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
self.wildcards_obj = None
|
||||
self.extranetwork_mappings_obj = None
|
||||
self.ppp_init = False
|
||||
self.ppp = None
|
||||
# log(self.ppp_logger, DEBUG_LEVEL.minimal, logging.INFO, f"Initializing {self.name} instance {self.instance_index}")
|
||||
|
||||
try:
|
||||
# Support for SD.Next
|
||||
import installer # type: ignore
|
||||
import installer # type: ignore # pylint: disable=import-outside-toplevel
|
||||
|
||||
if hasattr(installer, "control_extensions"):
|
||||
if self.title() not in installer.control_extensions:
|
||||
installer.control_extensions.append(self.title()) # We add the extension to the whitelist.
|
||||
@@ -86,7 +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.")
|
||||
pass
|
||||
except Exception as e: # pylint: disable=broad-except
|
||||
log(self.ppp_logger, DEBUG_LEVEL.minimal, logging.ERROR, f"Error while adding to the SD.Next extension whitelist: {e}")
|
||||
log(
|
||||
self.ppp_logger,
|
||||
DEBUG_LEVEL.minimal,
|
||||
logging.ERROR,
|
||||
f"Error while adding to the SD.Next extension whitelist: {e}",
|
||||
)
|
||||
|
||||
def title(self):
|
||||
"""
|
||||
@@ -113,7 +131,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
with gr.Accordion(PromptPostProcessor.NAME, open=False):
|
||||
force_equal_seeds = gr.Checkbox(
|
||||
label="Force equal seeds",
|
||||
info="Force all image seeds and variation seeds to be equal to the first one, disabling the default autoincrease.",
|
||||
info="Force all image seeds and variation seeds to be equal to the first one, disabling the default autoincrement.",
|
||||
value=False,
|
||||
# show_label=True,
|
||||
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.
|
||||
* Any other seed value and "Incremental seed" checked will use the specified seed for the first prompt and consecutive values for the rest.
|
||||
* Any other seed value and "Incremental seed" unchecked will use the specified seed for all the prompts.
|
||||
|
||||
Seeds are only used for the wildcards and choice constructs.
|
||||
""")
|
||||
gr.HTML("<br>")
|
||||
with gr.Row(equal_height=True):
|
||||
@@ -155,34 +171,51 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
elem_id="ppp_incremental_seed",
|
||||
)
|
||||
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):
|
||||
combinatorial = gr.Checkbox(
|
||||
label="Combinatorial mode",
|
||||
info="Generate all prompt combinations and cycle through them to fill the batch.",
|
||||
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,
|
||||
results_limit = gr.Number(
|
||||
label="Results limit (0 = no limit)",
|
||||
value=PromptPostProcessor.DEFAULT_RESULTS_LIMIT,
|
||||
precision=0,
|
||||
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 [
|
||||
force_equal_seeds,
|
||||
unlink_seed,
|
||||
seed,
|
||||
incremental_seed,
|
||||
combinatorial,
|
||||
combinatorial_shuffle,
|
||||
combinatorial_limit,
|
||||
run_mode,
|
||||
results_limit,
|
||||
results_shuffle,
|
||||
comb_random_fixed,
|
||||
default_sampler,
|
||||
]
|
||||
|
||||
def process(
|
||||
@@ -192,9 +225,11 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
input_unlink_seed,
|
||||
input_seed,
|
||||
input_incremental_seed,
|
||||
input_combinatorial,
|
||||
input_combinatorial_shuffle,
|
||||
input_combinatorial_limit,
|
||||
input_run_mode,
|
||||
input_results_limit,
|
||||
input_results_shuffle,
|
||||
input_comb_random_fixed,
|
||||
input_default_sampler,
|
||||
): # pylint: disable=arguments-differ
|
||||
"""
|
||||
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_seed (int): The seed value.
|
||||
input_incremental_seed (bool): Flag indicating whether to use incremental seed.
|
||||
input_combinatorial (bool): Flag indicating whether to use combinatorial mode.
|
||||
input_combinatorial_shuffle (bool): Flag indicating whether to shuffle the combinatorial results.
|
||||
input_combinatorial_limit (int): Maximum number of combinations (0 = no limit).
|
||||
input_run_mode (str): The run mode for prompt processing.
|
||||
input_results_limit (int): Maximum number of results (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:
|
||||
None
|
||||
@@ -267,10 +304,15 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
cup_remove_extranetwork_tags=getattr(
|
||||
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),
|
||||
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:
|
||||
self.ppp_init = True
|
||||
@@ -301,9 +343,16 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
"PPP unlink seed": input_unlink_seed,
|
||||
"PPP prompt seed": input_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(
|
||||
self.ppp_logger,
|
||||
@@ -311,17 +360,17 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
logging.INFO,
|
||||
f"Post-processing prompts ({'i2i' if is_i2i else 't2i'})",
|
||||
)
|
||||
env_info = {
|
||||
"app": app.value,
|
||||
"models_path": models_path,
|
||||
"model_filename": getattr(p.sd_model.sd_checkpoint_info, "filename", ""),
|
||||
"model_class": (
|
||||
env_info = PPPEnvInfo(
|
||||
app=app,
|
||||
models_path=models_path,
|
||||
model_filename=getattr(p.sd_model.sd_checkpoint_info, "filename", ""),
|
||||
model_class=(
|
||||
p.sd_model.model_config.__class__.__name__
|
||||
if app in (SUPPORTED_APPS.forge, SUPPORTED_APPS.forgeneo)
|
||||
else p.sd_model.__class__.__name__
|
||||
),
|
||||
"property_base": p.sd_model,
|
||||
}
|
||||
property_base=p.sd_model,
|
||||
)
|
||||
wc_wildcards_folders = getattr(opts, "ppp_wil_wildcardsfolders", "")
|
||||
if wc_wildcards_folders == "":
|
||||
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.extranetwork_mappings_obj.refresh_extranetwork_mappings(self.ppp_debug_level, enmappings_folders)
|
||||
ppp = PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
env_info,
|
||||
options,
|
||||
self.grammar_content,
|
||||
self.ppp_interrupt,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_mappings_obj,
|
||||
if self.ppp is None:
|
||||
self.ppp = PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
env_info,
|
||||
options,
|
||||
self.grammar_content,
|
||||
self.ppp_interrupt,
|
||||
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:
|
||||
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:
|
||||
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "Using unlinked 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)]
|
||||
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:
|
||||
calculated_seeds = [input_seed for _ in range(num_seeds)]
|
||||
else:
|
||||
@@ -388,8 +447,8 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
else:
|
||||
calculated_seeds = seeds
|
||||
|
||||
# (prompt type, typeindex) -> (new positive prompt, new negative prompt)
|
||||
prompts_list: dict[tuple[str, int], tuple[str, str]] = {}
|
||||
# [index][indextype] -> (new positive prompt, new negative prompt)
|
||||
prompts_list: list[list[tuple[str, str]]] = []
|
||||
extra_params = {}
|
||||
|
||||
# adds prompts
|
||||
@@ -402,24 +461,27 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
rnh: list[str] = getattr(p, "all_hr_negative_prompts", None)
|
||||
hiresfix_exists = bool(rph) and bool(rnh)
|
||||
for i in range(len(calculated_seeds)):
|
||||
if regular_exists:
|
||||
prompts_list[(regular_type, i)] = None
|
||||
prompts_list.append([])
|
||||
prompts_list[i].append(None)
|
||||
if hiresfix_exists:
|
||||
prompts_list[(hiresfix_type, i)] = None
|
||||
prompts_list[i].append(None)
|
||||
|
||||
ppp.process_prompts_group_start()
|
||||
if input_combinatorial:
|
||||
seed_for_comb = calculated_seeds[0] if calculated_seeds else 0
|
||||
self.ppp.process_prompts_group_start()
|
||||
if input_run_mode in (RUN_MODE.multiple.value, RUN_MODE.combinatorial.value):
|
||||
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)
|
||||
regular_changes = False
|
||||
hiresfix_changes = False
|
||||
if regular_exists:
|
||||
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "processing prompts combinatorially (regular)")
|
||||
comb_results = ppp.process_prompt(
|
||||
if input_run_mode == RUN_MODE.combinatorial.value:
|
||||
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],
|
||||
rnr[0],
|
||||
seed_for_comb,
|
||||
calculated_seeds,
|
||||
jobinfo={
|
||||
"job_timestamp": shared.state.job_timestamp,
|
||||
"job": shared.state.job,
|
||||
@@ -429,8 +491,12 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
num_comb = len(comb_results)
|
||||
for i in range(len(rpr)): # pylint: disable=consider-using-enumerate
|
||||
posp, negp, _ = comb_results[i % num_comb]
|
||||
prompts_list[(regular_type, i)] = (posp, negp)
|
||||
extra_params["PPP combination"] = [str(1 + (i % num_comb)) for i in range(len(rpr))]
|
||||
prompts_list[i][0] = (posp, negp)
|
||||
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:
|
||||
hiresfix_equal = regular_exists and rph == rpr and rnh == rnr
|
||||
if hiresfix_equal:
|
||||
@@ -438,21 +504,25 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
self.ppp_logger,
|
||||
self.ppp_debug_level,
|
||||
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
|
||||
prompts_list[(hiresfix_type, i)] = prompts_list.get((regular_type, i))
|
||||
prompts_list[i][1] = prompts_list[i][0]
|
||||
else:
|
||||
if input_run_mode == RUN_MODE.combinatorial.value:
|
||||
msg = "processing prompts combinatorially (hiresfix)"
|
||||
else:
|
||||
msg = "processing prompts for multiple results (hiresfix)"
|
||||
log(
|
||||
self.ppp_logger,
|
||||
self.ppp_debug_level,
|
||||
logging.INFO,
|
||||
"processing prompts combinatorially (hiresfix)",
|
||||
msg,
|
||||
)
|
||||
comb_results_hr = ppp.process_prompt(
|
||||
comb_results_hr = self.ppp.process_prompt(
|
||||
rph[0],
|
||||
rnh[0],
|
||||
seed_for_comb,
|
||||
calculated_seeds,
|
||||
jobinfo={
|
||||
"job_timestamp": shared.state.job_timestamp,
|
||||
"job": shared.state.job,
|
||||
@@ -462,64 +532,86 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
num_comb_hr = len(comb_results_hr)
|
||||
for i in range(len(rph)): # pylint: disable=consider-using-enumerate
|
||||
posp, negp, _ = comb_results_hr[i % num_comb_hr]
|
||||
prompts_list[(hiresfix_type, i)] = (posp, negp)
|
||||
extra_params["PPP HR combination"] = [str(1 + (i % num_comb_hr)) for i in range(len(rph))]
|
||||
prompts_list[i][1] = (posp, negp)
|
||||
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:
|
||||
# processes prompts
|
||||
for prompttype, typeindex in prompts_list.keys():
|
||||
log(
|
||||
self.ppp_logger,
|
||||
self.ppp_debug_level,
|
||||
logging.INFO,
|
||||
f"processing prompts ({prompttype}[{typeindex+1}])",
|
||||
)
|
||||
key = (
|
||||
(hash_fullenv, calculated_seeds[typeindex], rpr[typeindex], rnr[typeindex])
|
||||
if prompttype == regular_type
|
||||
else (hash_fullenv, calculated_seeds[typeindex], rph[typeindex], rnh[typeindex])
|
||||
)
|
||||
cached = self.lru_cache.get(key)
|
||||
if cached is None:
|
||||
hsh, seed, prompt, negative_prompt = key
|
||||
results = ppp.process_prompt(
|
||||
prompt,
|
||||
negative_prompt,
|
||||
seed,
|
||||
jobinfo={
|
||||
"job_timestamp": shared.state.job_timestamp,
|
||||
"job": shared.state.job,
|
||||
"detail": f"{prompttype} prompt",
|
||||
},
|
||||
for index, grouplist in enumerate(prompts_list):
|
||||
for typeindex in range(len(grouplist)):
|
||||
typeprompt = [regular_type, hiresfix_type][typeindex]
|
||||
log(
|
||||
self.ppp_logger,
|
||||
self.ppp_debug_level,
|
||||
logging.INFO,
|
||||
f"processing prompts ({typeprompt}[{index+1}])",
|
||||
)
|
||||
posp, negp, _ = results[0]
|
||||
cached = (posp, negp)
|
||||
self.lru_cache.put(key, cached)
|
||||
# adds also the result so i2i doesn't process it unnecessarily
|
||||
self.lru_cache.put((hsh, seed, posp, negp), cached)
|
||||
else:
|
||||
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "result already in cache")
|
||||
prompts_list[(prompttype, typeindex)] = cached
|
||||
ppp.process_prompts_group_end()
|
||||
key = (
|
||||
(hash_fullenv, calculated_seeds[index], rpr[index], rnr[index])
|
||||
if typeindex == 0
|
||||
else (hash_fullenv, calculated_seeds[index], rph[index], rnh[index])
|
||||
)
|
||||
cached = self.lru_cache.get(key)
|
||||
if cached is None:
|
||||
hsh, seed, prompt, negative_prompt = key
|
||||
if typeindex > 0:
|
||||
prev_prompts = prompts_list[index][typeindex - 1]
|
||||
input_vars = {
|
||||
"prev_pos_prompt": prev_prompts[0],
|
||||
"prev_neg_prompt": prev_prompts[1],
|
||||
}
|
||||
else:
|
||||
input_vars = None
|
||||
results = 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
|
||||
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)
|
||||
regular_changes = False
|
||||
hiresfix_changes = False
|
||||
for (prompttype, typeindex), (posp, negp) in prompts_list.items():
|
||||
if prompttype == regular_type:
|
||||
if rpr[typeindex].strip() != posp.strip() or rnr[typeindex].strip() != negp.strip():
|
||||
regular_changes = True
|
||||
rpr[typeindex] = posp
|
||||
rnr[typeindex] = negp
|
||||
elif prompttype == hiresfix_type:
|
||||
if rph[typeindex].strip() != posp.strip() or rnh[typeindex].strip() != negp.strip():
|
||||
hiresfix_changes = True
|
||||
rph[typeindex] = posp
|
||||
rnh[typeindex] = negp
|
||||
for index, grouplist in enumerate(prompts_list):
|
||||
for typeindex, groupprompts in enumerate(grouplist):
|
||||
if groupprompts is None:
|
||||
continue
|
||||
posp, negp = groupprompts
|
||||
if typeindex == 0:
|
||||
if rpr[index].strip() != posp.strip() or rnr[index].strip() != negp.strip():
|
||||
regular_changes = True
|
||||
rpr[index] = posp
|
||||
rnr[index] = negp
|
||||
elif typeindex == 1:
|
||||
if rph[index].strip() != posp.strip() or rnh[index].strip() != negp.strip():
|
||||
hiresfix_changes = True
|
||||
rph[index] = posp
|
||||
rnh[index] = negp
|
||||
|
||||
# initialize extra generation parameters
|
||||
if add_prompts:
|
||||
if hiresfix_exists:
|
||||
extra_params["PPP Hires prompt"] = rph
|
||||
extra_params["PPP Hires negative prompt"] = rnh
|
||||
if regular_changes:
|
||||
extra_params["PPP original prompts"] = regular_copy[0]
|
||||
extra_params["PPP original negative prompts"] = regular_copy[1]
|
||||
|
||||
+40
-30
@@ -1,4 +1,4 @@
|
||||
from dataclasses import replace
|
||||
from dataclasses import replace, make_dataclass
|
||||
import difflib
|
||||
import logging
|
||||
from pathlib import Path
|
||||
@@ -6,7 +6,16 @@ from typing import Any, NamedTuple, Optional
|
||||
import unittest
|
||||
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_wildcards import PPPWildcards # type: ignore
|
||||
from ppp import PromptPostProcessor # type: ignore
|
||||
@@ -74,19 +83,22 @@ class TestPromptPostProcessorBase(unittest.TestCase):
|
||||
cup_extranetwork_tags=True,
|
||||
cup_merge_attention=True,
|
||||
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 "",
|
||||
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.wildcards_obj = PPPWildcards(self.lf.log)
|
||||
self.extranetwork_maps_obj = PPPExtraNetworkMappings(self.lf.log)
|
||||
@@ -119,8 +131,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
|
||||
def init_ppp(
|
||||
self,
|
||||
ppp: Optional[str | PromptPostProcessor] = None,
|
||||
combinatorial: bool = False,
|
||||
combinatorial_limit: int = 0,
|
||||
**kwargs,
|
||||
) -> PromptPostProcessor:
|
||||
if isinstance(ppp, str):
|
||||
if ppp == "nocup":
|
||||
@@ -142,8 +153,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
|
||||
cup_ands_eol=False,
|
||||
cup_extranetwork_tags=False,
|
||||
cup_merge_attention=False,
|
||||
do_combinatorial=combinatorial,
|
||||
combinatorial_limit=combinatorial_limit,
|
||||
**kwargs,
|
||||
),
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -157,8 +167,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
|
||||
replace(
|
||||
self.defopts,
|
||||
strict_operators=False,
|
||||
do_combinatorial=combinatorial,
|
||||
combinatorial_limit=combinatorial_limit,
|
||||
**kwargs,
|
||||
),
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -173,8 +182,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
|
||||
self.def_env_info,
|
||||
replace(
|
||||
self.defopts,
|
||||
do_combinatorial=combinatorial,
|
||||
combinatorial_limit=combinatorial_limit,
|
||||
**kwargs,
|
||||
),
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -200,10 +208,9 @@ class TestPromptPostProcessorBase(unittest.TestCase):
|
||||
seed: int = 1,
|
||||
ppp: Optional[str | PromptPostProcessor] = None,
|
||||
interrupted: bool = False,
|
||||
combinatorial: bool = False,
|
||||
combinatorial_limit: int = 0,
|
||||
specific_wc_folders: Optional[list[Path]] = None,
|
||||
specific_em_folders: Optional[list[Path]] = None,
|
||||
input_vars: Optional[dict[str, Any]] = None,
|
||||
):
|
||||
"""
|
||||
Process the prompt and compare the results with the expected prompts.
|
||||
@@ -214,10 +221,9 @@ class TestPromptPostProcessorBase(unittest.TestCase):
|
||||
seed (int, optional): The seed value. Defaults to 1.
|
||||
ppp (Optional[str | PromptPostProcessor], optional): The PromptPostProcessor instance or type. Defaults to None.
|
||||
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_em_folders (Optional[list[Path]], optional): A list of specific extranetwork mapping folders to refresh. Defaults to None.
|
||||
input_vars (Optional[dict[str, Any]], optional): A dictionary of input variables. Defaults to None.
|
||||
|
||||
Returns:
|
||||
None
|
||||
@@ -232,13 +238,13 @@ class TestPromptPostProcessorBase(unittest.TestCase):
|
||||
DEBUG_LEVEL.full,
|
||||
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 = (
|
||||
[OutputTuple("", "", None)]
|
||||
if expected_output is None
|
||||
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
|
||||
errors = []
|
||||
the_obj.process_prompts_group_start()
|
||||
@@ -246,6 +252,8 @@ class TestPromptPostProcessorBase(unittest.TestCase):
|
||||
input_prompts.prompt,
|
||||
input_prompts.negative_prompt,
|
||||
seed,
|
||||
jobinfo={"test_case": self.id()},
|
||||
input_vars=input_vars,
|
||||
)
|
||||
the_obj.process_prompts_group_end()
|
||||
self.assertTrue(
|
||||
@@ -254,7 +262,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
|
||||
)
|
||||
if not self.interrupted and expected_output is not None:
|
||||
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:
|
||||
found = None
|
||||
for r_prompt, r_negative_prompt, r_variables in result:
|
||||
@@ -264,7 +272,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
|
||||
if not found:
|
||||
errors.extend(
|
||||
[
|
||||
"Combination not found in output",
|
||||
"Result not found in output",
|
||||
"Prompt:",
|
||||
out_prompt,
|
||||
"Negative Prompt:",
|
||||
@@ -286,7 +294,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
|
||||
if missing_vars or incorrect_vars:
|
||||
errors.extend(
|
||||
[
|
||||
"Combination found, but variables do not match",
|
||||
"Result found, but variables do not match",
|
||||
"Prompt:",
|
||||
out_prompt,
|
||||
"Negative Prompt:",
|
||||
@@ -311,6 +319,8 @@ class TestPromptPostProcessorBase(unittest.TestCase):
|
||||
input_prompts.prompt,
|
||||
input_prompts.negative_prompt,
|
||||
seed,
|
||||
jobinfo={"test_case": self.id()},
|
||||
input_vars=input_vars,
|
||||
)
|
||||
self.assertTrue(
|
||||
self.interrupted == interrupted,
|
||||
|
||||
+135
-29
@@ -1,6 +1,4 @@
|
||||
from dataclasses import replace
|
||||
|
||||
from ppp import PromptPostProcessor # type: ignore
|
||||
from ppp_classes import DEFAULT_SAMPLER, NEXT_SEED, RUN_MODE # type: ignore
|
||||
from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase
|
||||
|
||||
if __name__ == "__main__":
|
||||
@@ -22,7 +20,6 @@ class TestChoices(TestPromptPostProcessorBase):
|
||||
)
|
||||
|
||||
def test_ch_cyclical(self): # cyclical sampler cycles through all choices
|
||||
ppp_instance = self.init_ppp("nocup")
|
||||
self.process(
|
||||
InputTuple("the choices are: {@choice1|choice2|choice3}", ""),
|
||||
[
|
||||
@@ -31,11 +28,10 @@ class TestChoices(TestPromptPostProcessorBase):
|
||||
OutputTuple("the choices are: choice3", ""),
|
||||
OutputTuple("the choices are: choice1", ""), # cycles back
|
||||
],
|
||||
ppp=ppp_instance,
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
def test_ch_cyclical_multiple_constructs(self): # two independent @ constructs cycle together
|
||||
ppp_instance = self.init_ppp("nocup")
|
||||
self.process(
|
||||
InputTuple("{@a|b} {@c|d}", ""),
|
||||
[
|
||||
@@ -45,7 +41,7 @@ class TestChoices(TestPromptPostProcessorBase):
|
||||
OutputTuple("b d", ""),
|
||||
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
|
||||
@@ -67,7 +63,6 @@ class TestChoices(TestPromptPostProcessorBase):
|
||||
)
|
||||
|
||||
def test_ch_cyclical_mixed_samplers(self): # @ construct cycles while a ~ construct alongside is unaffected
|
||||
ppp_instance = self.init_ppp("nocup")
|
||||
self.process(
|
||||
InputTuple("{@a|b|c} {x|y}", ""),
|
||||
[
|
||||
@@ -76,7 +71,7 @@ class TestChoices(TestPromptPostProcessorBase):
|
||||
OutputTuple("c x", ""),
|
||||
OutputTuple("a y", ""), # @ cycles back
|
||||
],
|
||||
ppp=ppp_instance,
|
||||
ppp="nocup",
|
||||
)
|
||||
|
||||
def test_ch_choices_withcomments(self): # choices with comments and multiline
|
||||
@@ -103,6 +98,13 @@ class TestChoices(TestPromptPostProcessorBase):
|
||||
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
|
||||
self.process(
|
||||
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(
|
||||
InputTuple("<lora:test1:1><lora:test2:{0.2|0.5|0.7|1}>", ""),
|
||||
OutputTuple("", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
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,
|
||||
),
|
||||
ppp=self.init_ppp(None, cup_remove_extranetwork_tags=True),
|
||||
)
|
||||
|
||||
def test_ch_cmd_includewildcard(self):
|
||||
@@ -156,14 +147,129 @@ class TestChoices(TestPromptPostProcessorBase):
|
||||
|
||||
def test_ch_combinatorial(self):
|
||||
self.process(
|
||||
InputTuple("{choice1|choice2|choice3}, ${v:{option1|option2}}", ""),
|
||||
InputTuple("{choice1|choice2|choice3}, ${v:{option1|option2}}, {~a|b}", ""),
|
||||
[
|
||||
OutputTuple("choice1, option1", ""),
|
||||
OutputTuple("choice1, option2", ""),
|
||||
OutputTuple("choice2, option1", ""),
|
||||
OutputTuple("choice2, option2", ""),
|
||||
OutputTuple("choice3, option1", ""),
|
||||
OutputTuple("choice3, option2", "", {"v": "option2"}),
|
||||
OutputTuple("choice1, option1, a", ""),
|
||||
OutputTuple("choice1, option2, b", ""),
|
||||
OutputTuple("choice2, option1, a", ""),
|
||||
OutputTuple("choice2, option2, b", ""),
|
||||
OutputTuple("choice3, option1, b", ""),
|
||||
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
@@ -1,10 +1,7 @@
|
||||
import logging
|
||||
from dataclasses import replace
|
||||
|
||||
from ppp import PromptPostProcessor # type: ignore
|
||||
from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit("This script must not be run directly")
|
||||
|
||||
@@ -38,36 +35,17 @@ class TestCleanup(TestPromptPostProcessorBase):
|
||||
self.process(
|
||||
InputTuple("this is a <lora:test:1> test__yaml/wildcard7__", ""),
|
||||
OutputTuple("this is a test", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
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,
|
||||
),
|
||||
ppp=self.init_ppp(None, cup_remove_extranetwork_tags=True),
|
||||
)
|
||||
|
||||
def test_cl_dontremoveseparatorsoneol(self): # don't remove separators on eol
|
||||
self.process(
|
||||
InputTuple("this is a test,\nsecond line", ""),
|
||||
OutputTuple("this is a test,\nsecond line", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.def_env_info,
|
||||
replace(
|
||||
self.defopts,
|
||||
cup_extra_separators2=False,
|
||||
cup_extra_separators_include_eol=False,
|
||||
),
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
ppp=self.init_ppp(
|
||||
None,
|
||||
cup_extra_separators2=False,
|
||||
cup_extra_separators_include_eol=False,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -85,27 +63,19 @@ class TestCleanup(TestPromptPostProcessorBase):
|
||||
(d:0.9)""",
|
||||
"",
|
||||
),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.def_env_info,
|
||||
replace(
|
||||
self.defopts,
|
||||
cup_empty_constructs=False,
|
||||
cup_extra_separators=True,
|
||||
cup_extra_separators2=False,
|
||||
cup_extra_separators_include_eol=False,
|
||||
cup_extra_spaces=False,
|
||||
cup_breaks=False,
|
||||
cup_breaks_eol=False,
|
||||
cup_ands=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,
|
||||
ppp=self.init_ppp(
|
||||
None,
|
||||
cup_empty_constructs=False,
|
||||
cup_extra_separators=True,
|
||||
cup_extra_separators2=False,
|
||||
cup_extra_separators_include_eol=False,
|
||||
cup_extra_spaces=False,
|
||||
cup_breaks=False,
|
||||
cup_breaks_eol=False,
|
||||
cup_ands=False,
|
||||
cup_ands_eol=False,
|
||||
cup_extranetwork_tags=False,
|
||||
cup_merge_attention=False,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -191,7 +161,9 @@ class TestCleanup(TestPromptPostProcessorBase):
|
||||
"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):
|
||||
self.process(
|
||||
InputTuple(r"text with \(escaped unmatched\]", ""),
|
||||
|
||||
+81
-80
@@ -1,3 +1,4 @@
|
||||
from dataclasses import replace
|
||||
from ppp import PromptPostProcessor # type: ignore
|
||||
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)", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"attention": "parentheses"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"attention": "parentheses"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -42,10 +43,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("test1 test2 test3", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"attention": "disable"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"attention": "disable"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -63,10 +64,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"attention": "remove"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"attention": "remove"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -84,10 +85,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"attention": "error"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"attention": "error"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -106,10 +107,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("test1", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"scheduling": "before"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"scheduling": "before"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -127,10 +128,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("test2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"scheduling": "after"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"scheduling": "after"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -148,10 +149,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("test1 test3", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"scheduling": "first"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"scheduling": "first"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -169,10 +170,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"scheduling": "remove"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"scheduling": "remove"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -190,10 +191,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"scheduling": "error"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"scheduling": "error"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -212,10 +213,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("test1", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"alternation": "first"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"alternation": "first"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -233,10 +234,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"alternation": "remove"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"alternation": "remove"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -254,10 +255,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"alternation": "error"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"alternation": "error"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -276,10 +277,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("test1\ntest2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"and": "eol"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"and": "eol"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -297,10 +298,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("test1, test2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"and": "comma"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"and": "comma"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -318,10 +319,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("test1 test2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"and": "remove"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"and": "remove"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -339,10 +340,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"and": "error"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"and": "error"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -361,10 +362,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("test1\ntest2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"break": "eol"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"break": "eol"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -382,10 +383,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("test1, test2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"break": "comma"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"break": "comma"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -403,10 +404,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("test1 test2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"break": "remove"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"break": "remove"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -424,10 +425,10 @@ class TestHosts(TestPromptPostProcessorBase):
|
||||
OutputTuple("", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"ppp_config": {"hosts": {"tests": {"break": "error"}}},
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
ppp_config={"hosts": {"tests": {"break": "error"}}},
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
|
||||
+38
-68
@@ -1,5 +1,4 @@
|
||||
from dataclasses import replace
|
||||
|
||||
from ppp import PromptPostProcessor # type: ignore
|
||||
from ppp_classes import ONWARNING_CHOICES # type: ignore
|
||||
from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase
|
||||
@@ -57,18 +56,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
|
||||
"",
|
||||
),
|
||||
OutputTuple("", "", {"v1": ""}),
|
||||
ppp=PromptPostProcessor(
|
||||
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,
|
||||
),
|
||||
ppp=self.init_ppp(None, on_warning=ONWARNING_CHOICES.warn),
|
||||
)
|
||||
|
||||
# Variable in extranetworks
|
||||
@@ -690,18 +678,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
|
||||
"",
|
||||
),
|
||||
OutputTuple("not OK", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
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,
|
||||
),
|
||||
ppp=self.init_ppp(None, on_warning=ONWARNING_CHOICES.warn),
|
||||
)
|
||||
|
||||
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", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
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,
|
||||
),
|
||||
ppp=self.init_ppp(None, on_warning=ONWARNING_CHOICES.warn),
|
||||
)
|
||||
|
||||
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", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
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,
|
||||
ppp=self.init_ppp(None, on_warning=ONWARNING_CHOICES.warn),
|
||||
)
|
||||
|
||||
# Input variables
|
||||
|
||||
def test_input_variables(self):
|
||||
self.process(
|
||||
InputTuple(
|
||||
"${_input_prev_positive_prompt}, high quality",
|
||||
"",
|
||||
),
|
||||
OutputTuple(
|
||||
"this is a test, high quality",
|
||||
"",
|
||||
),
|
||||
input_vars={"prev_positive_prompt": "this is a test"},
|
||||
)
|
||||
|
||||
# Command tests
|
||||
@@ -801,10 +771,10 @@ class TestVarCommands(TestPromptPostProcessorBase):
|
||||
OutputTuple("this is PONY", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors",
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
model_filename="./webui/models/Stable-diffusion/ponymodel.safetensors",
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -1034,10 +1004,10 @@ class TestVarCommands(TestPromptPostProcessorBase):
|
||||
OutputTuple("<lora:lorapony:0.8>inlinetrigger, triggerpony1, triggerpony2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors",
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
model_filename="./webui/models/Stable-diffusion/ponymodel.safetensors",
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -1055,10 +1025,10 @@ class TestVarCommands(TestPromptPostProcessorBase):
|
||||
OutputTuple("<lora:lorapony:0.4>inlinetrigger, triggerpony1, triggerpony2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors",
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
model_filename="./webui/models/Stable-diffusion/ponymodel.safetensors",
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -1076,10 +1046,10 @@ class TestVarCommands(TestPromptPostProcessorBase):
|
||||
OutputTuple("<lora:lorapony:0.6:0.8>inlinetrigger, triggerpony1, triggerpony2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors",
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
model_filename="./webui/models/Stable-diffusion/ponymodel.safetensors",
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -1097,10 +1067,10 @@ class TestVarCommands(TestPromptPostProcessorBase):
|
||||
OutputTuple("<lora:loraillustrious:0.9:0.8>inlinetrigger, triggerillustrious1, triggerillustrious2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"model_filename": "./webui/models/Stable-diffusion/ilxlmodel.safetensors",
|
||||
},
|
||||
replace(
|
||||
self.def_env_info,
|
||||
model_filename="./webui/models/Stable-diffusion/ilxlmodel.safetensors",
|
||||
),
|
||||
self.defopts,
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
|
||||
+10
-10
@@ -24,10 +24,10 @@ class TestModelVariants(TestPromptPostProcessorBase):
|
||||
OutputTuple("test1test2", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"model_filename": "./webui/models/Stable-diffusion/testmodel.safetensors",
|
||||
"ppp_config": {
|
||||
replace(
|
||||
self.def_env_info,
|
||||
model_filename="./webui/models/Stable-diffusion/testmodel.safetensors",
|
||||
ppp_config={
|
||||
"models": {
|
||||
"sd1": {
|
||||
"detect": {"tests": {"class": ["SD15", "SD15_instructpix2pix"]}},
|
||||
@@ -62,7 +62,7 @@ class TestModelVariants(TestPromptPostProcessorBase):
|
||||
},
|
||||
}
|
||||
},
|
||||
},
|
||||
),
|
||||
replace(
|
||||
self.defopts,
|
||||
on_warning=ONWARNING_CHOICES.warn,
|
||||
@@ -84,15 +84,15 @@ class TestModelVariants(TestPromptPostProcessorBase):
|
||||
OutputTuple("not SDXL, not PONY", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
{
|
||||
**self.def_env_info,
|
||||
"model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors",
|
||||
"ppp_config": {
|
||||
replace(
|
||||
self.def_env_info,
|
||||
model_filename="./webui/models/Stable-diffusion/ponymodel.safetensors",
|
||||
ppp_config={
|
||||
"models": {
|
||||
"sdxl": None,
|
||||
}
|
||||
},
|
||||
},
|
||||
),
|
||||
replace(
|
||||
self.defopts,
|
||||
on_warning=ONWARNING_CHOICES.warn,
|
||||
|
||||
+54
-83
@@ -1,7 +1,6 @@
|
||||
from dataclasses import replace
|
||||
|
||||
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
|
||||
|
||||
if __name__ == "__main__":
|
||||
@@ -19,18 +18,10 @@ class TestWildcards(TestPromptPostProcessorBase):
|
||||
self.process(
|
||||
InputTuple("__bad_wildcard__", "{option1|option2}"),
|
||||
OutputTuple("__bad_wildcard__", "{option1|option2}"),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.def_env_info,
|
||||
replace(
|
||||
self.defopts,
|
||||
process_wildcards=False,
|
||||
if_wildcards=IFWILDCARDS_CHOICES.ignore,
|
||||
),
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
ppp=self.init_ppp(
|
||||
None,
|
||||
process_wildcards=False,
|
||||
if_wildcards=IFWILDCARDS_CHOICES.ignore,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -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>",
|
||||
"[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]",
|
||||
),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.def_env_info,
|
||||
replace(
|
||||
self.defopts,
|
||||
process_wildcards=False,
|
||||
if_wildcards=IFWILDCARDS_CHOICES.remove,
|
||||
),
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
ppp=self.init_ppp(
|
||||
None,
|
||||
process_wildcards=False,
|
||||
if_wildcards=IFWILDCARDS_CHOICES.remove,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -63,18 +46,10 @@ class TestWildcards(TestPromptPostProcessorBase):
|
||||
self.process(
|
||||
InputTuple("__bad_wildcard__", "{option1|option2}"),
|
||||
OutputTuple(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", "{option1|option2}"),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.def_env_info,
|
||||
replace(
|
||||
self.defopts,
|
||||
process_wildcards=False,
|
||||
if_wildcards=IFWILDCARDS_CHOICES.warn,
|
||||
),
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
ppp=self.init_ppp(
|
||||
None,
|
||||
process_wildcards=False,
|
||||
if_wildcards=IFWILDCARDS_CHOICES.warn,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -85,18 +60,10 @@ class TestWildcards(TestPromptPostProcessorBase):
|
||||
PromptPostProcessor.WILDCARD_STOP.format("__bad_wildcard__") + "__bad_wildcard__",
|
||||
"{option1|option2}",
|
||||
),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.def_env_info,
|
||||
replace(
|
||||
self.defopts,
|
||||
process_wildcards=False,
|
||||
if_wildcards=IFWILDCARDS_CHOICES.stop,
|
||||
),
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
ppp=self.init_ppp(
|
||||
None,
|
||||
process_wildcards=False,
|
||||
if_wildcards=IFWILDCARDS_CHOICES.stop,
|
||||
),
|
||||
interrupted=True,
|
||||
)
|
||||
@@ -105,18 +72,10 @@ class TestWildcards(TestPromptPostProcessorBase):
|
||||
self.process(
|
||||
InputTuple("${v=__bad_wildcard__}${v}", ""),
|
||||
OutputTuple(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", ""),
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.def_env_info,
|
||||
replace(
|
||||
self.defopts,
|
||||
process_wildcards=False,
|
||||
if_wildcards=IFWILDCARDS_CHOICES.warn,
|
||||
),
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
ppp=self.init_ppp(
|
||||
None,
|
||||
process_wildcards=False,
|
||||
if_wildcards=IFWILDCARDS_CHOICES.warn,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -128,6 +87,22 @@ class TestWildcards(TestPromptPostProcessorBase):
|
||||
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
|
||||
self.process(
|
||||
InputTuple("the choices are: __text/wildcard1__", ""),
|
||||
@@ -337,6 +312,13 @@ class TestWildcards(TestPromptPostProcessorBase):
|
||||
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
|
||||
self.process(
|
||||
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, option2", "", {"v": "option2"}),
|
||||
],
|
||||
combinatorial=True,
|
||||
ppp=self.init_ppp(None, run_mode=RUN_MODE.combinatorial),
|
||||
)
|
||||
|
||||
def test_wc_combinatorial_2(self): # combinatorial wildcard
|
||||
@@ -545,8 +527,7 @@ class TestWildcards(TestPromptPostProcessorBase):
|
||||
OutputTuple("choice1-choice3", ""),
|
||||
OutputTuple("choice3-choice1", ""),
|
||||
],
|
||||
ppp="nocup",
|
||||
combinatorial=True,
|
||||
ppp=self.init_ppp("nocup", run_mode=RUN_MODE.combinatorial),
|
||||
)
|
||||
|
||||
def test_wc_combinatorial_3(self): # combinatorial wildcard (keep choice order)
|
||||
@@ -564,19 +545,11 @@ class TestWildcards(TestPromptPostProcessorBase):
|
||||
## choices 1 and 3
|
||||
OutputTuple("choice1-choice3", ""),
|
||||
],
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.def_env_info,
|
||||
replace(
|
||||
self.defopts,
|
||||
keep_choices_order=True,
|
||||
cup_do_cleanup=False,
|
||||
do_combinatorial=True,
|
||||
),
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
ppp=self.init_ppp(
|
||||
None,
|
||||
keep_choices_order=True,
|
||||
cup_do_cleanup=False,
|
||||
run_mode=RUN_MODE.combinatorial,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -603,8 +576,7 @@ class TestWildcards(TestPromptPostProcessorBase):
|
||||
OutputTuple("choice1-choice3", ""),
|
||||
OutputTuple("choice3-choice1", ""),
|
||||
],
|
||||
ppp="nocup",
|
||||
combinatorial=True,
|
||||
ppp=self.init_ppp("nocup", run_mode=RUN_MODE.combinatorial),
|
||||
)
|
||||
|
||||
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:loraany2:1> trigger3, trigger4, ", ""),
|
||||
],
|
||||
ppp="nocup",
|
||||
combinatorial=True,
|
||||
ppp=self.init_ppp("nocup", run_mode=RUN_MODE.combinatorial),
|
||||
)
|
||||
|
||||
@@ -6,7 +6,7 @@ yaml:
|
||||
- choice3
|
||||
|
||||
wildcard2:
|
||||
- ~r2-3'Wildcard description'$$-$$
|
||||
- r2-3'Wildcard description'$$-$$
|
||||
- "'label1,label2'4::choice1"
|
||||
- "3:: choice2 "
|
||||
- { labels: ["label1", "label3"], weight: 2, content: choice3 }
|
||||
@@ -111,6 +111,14 @@ yaml:
|
||||
- if _sd in ("test1", "test2")::4
|
||||
- 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:
|
||||
- 5::__yaml/circular2__
|
||||
- choice1
|
||||
|
||||
@@ -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()
|
||||
@@ -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.
|
||||
@@ -15,14 +15,13 @@ Main PPP node that processes prompts.
|
||||
* **process_wildcards**: Activates the wildcard processing.
|
||||
* **do_cleanup**: Activates the cleanup processing.
|
||||
* **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.
|
||||
* **combinatorial_shuffle**: It shuffles the combinatorial results.
|
||||
* **combinatorial_limit**: Limit for the number of generated combinations.
|
||||
* **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.
|
||||
* **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.
|
||||
* **wc_options**: Connection to a Wildcards options node.
|
||||
* **stn_options**: Connection to a Send-To-Negative options node.
|
||||
* **cup_options**: Connection to a Cleanup 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.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user