* Converted do_combinatorial property to run_mode, with a new multiple mode. This changes the node properties in ComfyUI.
* Added option to set the default choice sampler. * ComfyUI: Added ACBPPPRunModeOptions node to set run mode options. * A1111: ppp object is now kept between executions, so cyclical state is saved. * Added documentation in the cookbook regarding seed behavior and cyclical sampler resets. * Adjusted the testing methods for more flexibility. * Some refactoring.
This commit is contained in:
@@ -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 — do not hardcode option dicts from scratch.
|
||||
|
||||
```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,
|
||||
|
||||
@@ -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",
|
||||
|
||||
+31
-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,15 @@ 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.
|
||||
|
||||
### ACB PPP Wildcard Options node
|
||||
|
||||
Options for wildcard processing, in case you want to change them from the defaults.
|
||||
@@ -149,9 +169,11 @@ 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.
|
||||
|
||||
### General settings
|
||||
|
||||
|
||||
@@ -570,6 +570,40 @@ Relative paths are resolved against the `logs` folder inside the extension direc
|
||||
> [!TIP]
|
||||
> The `.jsonl` format is the most convenient for programmatic processing. The `.yaml` format is the easiest to read manually.
|
||||
|
||||
## Seed behavior per host
|
||||
|
||||
The seed determines which choices are picked for wildcards and `~` (random) samplers. Each host handles it differently.
|
||||
|
||||
### A1111 and derivatives
|
||||
|
||||
By default, PPP uses each image's own seed, which A1111 auto-increments across the batch. Each image therefore gets independently seeded wildcard expansions.
|
||||
|
||||
The extension provides options to change this:
|
||||
|
||||
- **Force equal seeds**: sets every image 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, all result variants are generated once using the first seed in the batch. The images then cycle through those pre-generated results in order; no additional seed is used for subsequent images.
|
||||
|
||||
### ComfyUI
|
||||
|
||||
The seed is an explicit node input (default: -1 for random). Each node execution uses exactly the provided seed. There is no batch; each execution processes one prompt independently.
|
||||
|
||||
## 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.
|
||||
|
||||
@@ -21,6 +21,7 @@ from ppp_classes import (
|
||||
ModelConfig,
|
||||
ModelDetectConfig,
|
||||
PPPException,
|
||||
RUN_MODE,
|
||||
VariantConfig,
|
||||
PPPConfig,
|
||||
IFWILDCARDS_CHOICES,
|
||||
@@ -78,10 +79,11 @@ 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_COMBINATORIAL_RANDOMSAMPLER_FIXED = defopt["combinatorial_randomsampler_fixed"]
|
||||
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"]
|
||||
|
||||
WILDCARD_WARNING = '(WARNING TEXT "INVALID WILDCARD" IN BRIGHT RED:1.5)\nBREAK '
|
||||
@@ -309,7 +311,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
if user_config_file == "":
|
||||
if app == SUPPORTED_APPS.comfyui.value:
|
||||
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():
|
||||
@@ -1126,21 +1128,27 @@ 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}", DEBUG_LEVEL.minimal)
|
||||
self.log(logging.INFO, f"Combinatorial: {self.state.options.do_combinatorial}")
|
||||
self.log(logging.INFO, f"Run mode: {self.state.options.run_mode.name}")
|
||||
if self.state.options.run_mode == RUN_MODE.combinatorial:
|
||||
self.log(logging.INFO, f"Up to {self.state.options.results_limit} combinations")
|
||||
elif self.state.options.run_mode == RUN_MODE.multiple:
|
||||
self.log(logging.INFO, f"Up to {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."""
|
||||
|
||||
+16
-4
@@ -47,6 +47,17 @@ class ONWARNING_CHOICES(Enum):
|
||||
stop = "stop"
|
||||
|
||||
|
||||
class RUN_MODE(Enum):
|
||||
single = "single"
|
||||
multiple = "multiple"
|
||||
combinatorial = "combinatorial"
|
||||
|
||||
|
||||
class DEFAULT_SAMPLER(Enum):
|
||||
random = "random"
|
||||
cyclical = "cyclical"
|
||||
|
||||
|
||||
# ------------------- Host configuration -------------------
|
||||
|
||||
AttentionOption = Literal["ok", "parentheses", "disable", "remove", "error"]
|
||||
@@ -209,11 +220,12 @@ 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
|
||||
combinatorial_randomsampler_fixed: bool = True # if True, the random sampler will be fixed across all DFS runs
|
||||
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
|
||||
|
||||
def __post_init__(self):
|
||||
if not self.cup_do_cleanup:
|
||||
|
||||
+120
-52
@@ -1,17 +1,26 @@
|
||||
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, PPPException, PPPStateOptions
|
||||
from ppp_classes import (
|
||||
DEFAULT_SAMPLER,
|
||||
IFWILDCARDS_CHOICES,
|
||||
ONWARNING_CHOICES,
|
||||
SUPPORTED_APPS,
|
||||
PPPException,
|
||||
RUN_MODE,
|
||||
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,38 +199,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",
|
||||
},
|
||||
),
|
||||
"combinatorial_randomsampler_fixed": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": PromptPostProcessor.DEFAULT_COMBINATORIAL_RANDOMSAMPLER_FIXED,
|
||||
"tooltip": "Fix the value of any specified random samplers across all combinations in combinatorial mode",
|
||||
"label_on": "Yes",
|
||||
"label_off": "No",
|
||||
"options": [e.value for e in RUN_MODE],
|
||||
"default": PromptPostProcessor.DEFAULT_RUN_MODE,
|
||||
"tooltip": "Run mode",
|
||||
},
|
||||
),
|
||||
"wc_options": (
|
||||
@@ -252,12 +243,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",
|
||||
},
|
||||
),
|
||||
},
|
||||
@@ -292,9 +282,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"
|
||||
@@ -313,20 +303,18 @@ class PromptPostProcessorComfyUINode:
|
||||
seed,
|
||||
debug_level,
|
||||
on_warnings,
|
||||
strict_operators,
|
||||
process_wildcards,
|
||||
do_cleanup,
|
||||
cleanup_variables,
|
||||
do_combinatorial,
|
||||
combinatorial_shuffle,
|
||||
combinatorial_limit,
|
||||
combinatorial_randomsampler_fixed,
|
||||
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
|
||||
@@ -377,12 +365,14 @@ class PromptPostProcessorComfyUINode:
|
||||
|
||||
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
|
||||
),
|
||||
@@ -433,11 +423,16 @@ class PromptPostProcessorComfyUINode:
|
||||
if cup_options
|
||||
else PromptPostProcessor.DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS
|
||||
),
|
||||
do_combinatorial=do_combinatorial,
|
||||
combinatorial_shuffle=combinatorial_shuffle,
|
||||
combinatorial_limit=combinatorial_limit,
|
||||
combinatorial_randomsampler_fixed=combinatorial_randomsampler_fixed,
|
||||
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
|
||||
),
|
||||
)
|
||||
self.wildcards_obj.refresh_wildcards(
|
||||
options.debug_level,
|
||||
@@ -481,6 +476,78 @@ class PromptPostProcessorComfyUINode:
|
||||
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",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
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,
|
||||
):
|
||||
options = {
|
||||
"results_limit": results_limit,
|
||||
"results_shuffle": results_shuffle,
|
||||
"comb_random_fixed": comb_random_fixed,
|
||||
"default_sampler": default_sampler,
|
||||
}
|
||||
return (options,)
|
||||
|
||||
|
||||
class PromptPostProcessorWildcardOptionsComfyUINode:
|
||||
"""
|
||||
Node for wildcard options.
|
||||
@@ -1050,6 +1117,7 @@ class PromptPostProcessorWildcardConcatComfyUINode:
|
||||
|
||||
|
||||
try:
|
||||
# pylint: disable=import-error
|
||||
from server import PromptServer # type: ignore
|
||||
from aiohttp import web as _aiohttp_web # type: ignore
|
||||
|
||||
|
||||
@@ -216,6 +216,7 @@ def get_model_config_from_filename(filename: Path) -> object | None:
|
||||
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
|
||||
|
||||
+92
-47
@@ -11,7 +11,7 @@ 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, 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
|
||||
@@ -88,7 +88,7 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
)
|
||||
|
||||
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 = []
|
||||
@@ -116,33 +116,59 @@ 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
|
||||
)
|
||||
|
||||
if not self.state.options.do_combinatorial:
|
||||
self.__forced_path = list(self.state.cyclical_state.current_path)
|
||||
self.__trace = []
|
||||
self.visit(parsed)
|
||||
self.__finalize_variables()
|
||||
if self.__trace:
|
||||
self.state.cyclical_state.last_trace = self.__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.__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()))
|
||||
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; results will start repeating.")
|
||||
warned_cycle = True
|
||||
return results
|
||||
|
||||
# Combinatorial mode: explore every possible path through choices and wildcards via DFS.
|
||||
# __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.
|
||||
# __rand_decisions caches random (~) choices so they stay consistent across all runs.
|
||||
initial_vars = self.state.variables.backup_user()
|
||||
results: list[tuple[str, list[tuple[str, bool]], tuple]] = []
|
||||
limit = self.state.options.combinatorial_limit
|
||||
self.__rand_decisions = {}
|
||||
|
||||
def _run(forced_path: tuple[int, ...]) -> tuple[int, ...]:
|
||||
@@ -165,7 +191,7 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
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)
|
||||
@@ -174,12 +200,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.
|
||||
@@ -188,7 +214,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):
|
||||
@@ -1647,8 +1675,18 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
else_mapping = v
|
||||
num_mappings = len(found_mappings)
|
||||
if num_mappings > 0:
|
||||
if self.state.options.do_combinatorial:
|
||||
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.__forced_path[decision_idx], num_mappings - 1)
|
||||
@@ -1656,15 +1694,19 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
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,
|
||||
@@ -1878,7 +1920,11 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
if options is None:
|
||||
options = {}
|
||||
specified_sampler = options.get("sampler")
|
||||
sampler: str = specified_sampler if specified_sampler is not None else "~"
|
||||
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:
|
||||
@@ -1932,7 +1978,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)
|
||||
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,
|
||||
@@ -1952,17 +1998,14 @@ class TreeProcessor(lark.visitors.Interpreter):
|
||||
else:
|
||||
all_selections.extend(permutations(available_choices, k))
|
||||
decision_idx = len(self.__trace)
|
||||
if self.state.options.do_combinatorial:
|
||||
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.combinatorial_randomsampler_fixed
|
||||
):
|
||||
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 combinatorial_randomsampler_fixed is True, we only do this once per decision
|
||||
# index, so all combinations share the same choice.
|
||||
# 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)]]
|
||||
@@ -2051,11 +2094,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
|
||||
|
||||
@@ -2397,7 +2441,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)}'")
|
||||
|
||||
+119
-66
@@ -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,22 @@ 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,
|
||||
SUPPORTED_APPS,
|
||||
SUPPORTED_APPS_NAMES,
|
||||
RUN_MODE,
|
||||
PPPStateOptions,
|
||||
)
|
||||
from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory, log
|
||||
from ppp_cache import PPPLRUCache
|
||||
from ppp_wildcards import PPPWildcards
|
||||
@@ -74,11 +83,12 @@ 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:
|
||||
@@ -161,41 +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",
|
||||
)
|
||||
combinatorial_randomsampler_fixed = gr.Checkbox(
|
||||
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_COMBINATORIAL_RANDOMSAMPLER_FIXED,
|
||||
elem_id="ppp_combinatorial_randomsampler_fixed",
|
||||
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,
|
||||
combinatorial_randomsampler_fixed,
|
||||
run_mode,
|
||||
results_limit,
|
||||
results_shuffle,
|
||||
comb_random_fixed,
|
||||
default_sampler,
|
||||
]
|
||||
|
||||
def process(
|
||||
@@ -205,10 +225,11 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
input_unlink_seed,
|
||||
input_seed,
|
||||
input_incremental_seed,
|
||||
input_combinatorial,
|
||||
input_combinatorial_shuffle,
|
||||
input_combinatorial_limit,
|
||||
input_combinatorial_randomsampler_fixed,
|
||||
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.
|
||||
@@ -219,10 +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_combinatorial_randomsampler_fixed (bool): Flag indicating whether to fix the random sampler across all combinations.
|
||||
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
|
||||
@@ -282,11 +304,12 @@ 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,
|
||||
combinatorial_randomsampler_fixed=input_combinatorial_randomsampler_fixed,
|
||||
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)
|
||||
)
|
||||
if not self.ppp_init:
|
||||
self.ppp_init = True
|
||||
@@ -317,10 +340,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 combinatorial random sampler fixed": input_combinatorial_randomsampler_fixed,
|
||||
"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,
|
||||
@@ -362,16 +391,24 @@ 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,
|
||||
)
|
||||
hash_fullenv = hash((ppp.envinfo_hash, ppp.options_hash, 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))
|
||||
|
||||
if input_force_equal_seeds:
|
||||
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "Forcing equal seeds")
|
||||
@@ -424,16 +461,20 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
if hiresfix_exists:
|
||||
prompts_list[i].append(None)
|
||||
|
||||
ppp.process_prompts_group_start()
|
||||
if input_combinatorial:
|
||||
self.ppp.process_prompts_group_start()
|
||||
if input_run_mode in (RUN_MODE.multiple.value, RUN_MODE.combinatorial.value):
|
||||
seed_for_comb = calculated_seeds[0] if calculated_seeds else 0
|
||||
regular_copy = (rpr.copy() if rpr else None, rnr.copy() if rnr else None)
|
||||
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,
|
||||
@@ -447,7 +488,11 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
for i in range(len(rpr)): # pylint: disable=consider-using-enumerate
|
||||
posp, negp, _ = comb_results[i % num_comb]
|
||||
prompts_list[i][0] = (posp, negp)
|
||||
extra_params["PPP combination"] = [str(1 + (i % num_comb)) for i in range(len(rpr))]
|
||||
if input_run_mode == RUN_MODE.combinatorial.value:
|
||||
field_name = "PPP combination"
|
||||
else:
|
||||
field_name = "PPP result"
|
||||
extra_params[field_name] = [str(1 + (i % num_comb)) for i in range(len(rpr))]
|
||||
if hiresfix_exists:
|
||||
hiresfix_equal = regular_exists and rph == rpr and rnh == rnr
|
||||
if hiresfix_equal:
|
||||
@@ -455,18 +500,22 @@ 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[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,
|
||||
@@ -480,7 +529,11 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
for i in range(len(rph)): # pylint: disable=consider-using-enumerate
|
||||
posp, negp, _ = comb_results_hr[i % num_comb_hr]
|
||||
prompts_list[i][1] = (posp, negp)
|
||||
extra_params["PPP HR combination"] = [str(1 + (i % num_comb_hr)) for i in range(len(rph))]
|
||||
if input_run_mode == RUN_MODE.combinatorial.value:
|
||||
field_name = "PPP HR combination"
|
||||
else:
|
||||
field_name = "PPP HR result"
|
||||
extra_params[field_name] = [str(1 + (i % num_comb_hr)) for i in range(len(rph))]
|
||||
else:
|
||||
# processes prompts
|
||||
for index, grouplist in enumerate(prompts_list):
|
||||
@@ -508,7 +561,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
}
|
||||
else:
|
||||
input_vars = None
|
||||
results = ppp.process_prompt(
|
||||
results = self.ppp.process_prompt(
|
||||
prompt,
|
||||
negative_prompt,
|
||||
seed,
|
||||
@@ -527,7 +580,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
else:
|
||||
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "result already in cache")
|
||||
prompts_list[index][typeindex] = cached
|
||||
ppp.process_prompts_group_end()
|
||||
self.ppp.process_prompts_group_end()
|
||||
|
||||
# updates the prompts
|
||||
regular_copy = (rpr.copy() if rpr else None, rnr.copy() if rnr else None)
|
||||
@@ -538,7 +591,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
|
||||
for typeindex, groupprompts in enumerate(grouplist):
|
||||
if groupprompts is None:
|
||||
continue
|
||||
(posp, negp) = groupprompts
|
||||
posp, negp = groupprompts
|
||||
if typeindex == 0:
|
||||
if rpr[index].strip() != posp.strip() or rnr[index].strip() != negp.strip():
|
||||
regular_changes = True
|
||||
|
||||
+15
-22
@@ -6,7 +6,7 @@ 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, ONWARNING_CHOICES, RUN_MODE, PPPStateOptions
|
||||
from ppp_enmappings import PPPExtraNetworkMappings # type: ignore
|
||||
from ppp_wildcards import PPPWildcards # type: ignore
|
||||
from ppp import PromptPostProcessor # type: ignore
|
||||
@@ -74,11 +74,12 @@ 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,
|
||||
combinatorial_randomsampler_fixed=True,
|
||||
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,
|
||||
)
|
||||
self.def_env_info = {
|
||||
"app": "tests",
|
||||
@@ -120,8 +121,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":
|
||||
@@ -143,8 +143,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,
|
||||
@@ -158,8 +157,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
|
||||
replace(
|
||||
self.defopts,
|
||||
strict_operators=False,
|
||||
do_combinatorial=combinatorial,
|
||||
combinatorial_limit=combinatorial_limit,
|
||||
**kwargs,
|
||||
),
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
@@ -174,8 +172,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,
|
||||
@@ -201,8 +198,6 @@ 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,
|
||||
@@ -216,8 +211,6 @@ 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.
|
||||
@@ -235,13 +228,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()
|
||||
@@ -259,7 +252,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:
|
||||
@@ -269,7 +262,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:",
|
||||
@@ -291,7 +284,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:",
|
||||
|
||||
+42
-35
@@ -1,6 +1,4 @@
|
||||
from dataclasses import replace
|
||||
|
||||
from ppp import PromptPostProcessor # type: ignore
|
||||
from ppp_classes import DEFAULT_SAMPLER, 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
|
||||
@@ -138,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):
|
||||
@@ -172,23 +156,15 @@ class TestChoices(TestPromptPostProcessorBase):
|
||||
OutputTuple("choice3, option1, a", ""),
|
||||
OutputTuple("choice3, option2, a", "", {"v": "option2"}),
|
||||
],
|
||||
ppp=PromptPostProcessor(
|
||||
self.ppp_logger,
|
||||
self.def_env_info,
|
||||
replace(
|
||||
self.defopts,
|
||||
do_combinatorial=True,
|
||||
combinatorial_randomsampler_fixed=False, # allow different random choices across combinations
|
||||
),
|
||||
self.grammar_content,
|
||||
self.interrupt,
|
||||
self.wildcards_obj,
|
||||
self.extranetwork_maps_obj,
|
||||
ppp=self.init_ppp(
|
||||
None,
|
||||
run_mode=RUN_MODE.combinatorial,
|
||||
comb_random_fixed=False, # allow different random choices across combinations
|
||||
),
|
||||
)
|
||||
|
||||
def test_ch_combinatorial_random_consistent(self): # ~ sampler picks one value shared across all combinations
|
||||
ppp_instance = self.init_ppp("nocup", combinatorial=True)
|
||||
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}", "", seed=1)
|
||||
ppp_instance.process_prompts_group_end()
|
||||
@@ -199,3 +175,34 @@ class TestChoices(TestPromptPostProcessorBase):
|
||||
1,
|
||||
f"The ~ sampler must yield the same value across all combinations, got: {rnd_choices}",
|
||||
)
|
||||
|
||||
# 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,
|
||||
),
|
||||
)
|
||||
|
||||
+18
-47
@@ -1,7 +1,5 @@
|
||||
import logging
|
||||
from dataclasses import replace
|
||||
|
||||
from ppp import PromptPostProcessor
|
||||
from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase
|
||||
|
||||
if __name__ == "__main__":
|
||||
@@ -37,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,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -84,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,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
+4
-50
@@ -1,5 +1,3 @@
|
||||
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 +55,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 +677,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 +697,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 +717,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),
|
||||
)
|
||||
|
||||
# Input variables
|
||||
|
||||
+30
-83
@@ -1,7 +1,5 @@
|
||||
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
|
||||
|
||||
@@ -20,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,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -45,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,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -64,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,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -86,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,
|
||||
)
|
||||
@@ -106,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,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -516,7 +474,7 @@ inv.alid2:
|
||||
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
|
||||
@@ -569,8 +527,7 @@ inv.alid2:
|
||||
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)
|
||||
@@ -588,19 +545,11 @@ inv.alid2:
|
||||
## 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,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -627,8 +576,7 @@ inv.alid2:
|
||||
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
|
||||
@@ -650,6 +598,5 @@ inv.alid2:
|
||||
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),
|
||||
)
|
||||
|
||||
@@ -0,0 +1,14 @@
|
||||
# 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.
|
||||
|
||||
## 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