From 90b4e2a791c003c048a9fe3c5d1e7e6a2925e065 Mon Sep 17 00:00:00 2001 From: Antonio Cordero Balcazar Date: Tue, 21 Jul 2026 19:39:25 +0200 Subject: [PATCH] * Added support for a fixed random sampler across combinations in combinatorial mode. --- ppp.py | 1 + ppp_classes.py | 1 + ppp_comfyui.py | 11 ++++ ppp_tree.py | 128 +++++++++++++++++++++-------------------- scripts/ppp_script.py | 11 ++++ tests/base_tests.py | 1 + tests/tests_choices.py | 27 ++++++++- 7 files changed, 118 insertions(+), 62 deletions(-) diff --git a/ppp.py b/ppp.py index 282c2b0..d65c9a9 100644 --- a/ppp.py +++ b/ppp.py @@ -81,6 +81,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in 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_RESULTS_FILE = defopt["results_file"] WILDCARD_WARNING = '(WARNING TEXT "INVALID WILDCARD" IN BRIGHT RED:1.5)\nBREAK ' diff --git a/ppp_classes.py b/ppp_classes.py index 40da611..159cc0a 100644 --- a/ppp_classes.py +++ b/ppp_classes.py @@ -212,6 +212,7 @@ class PPPStateOptions: 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 def __post_init__(self): diff --git a/ppp_comfyui.py b/ppp_comfyui.py index cc8047a..11e5e3f 100644 --- a/ppp_comfyui.py +++ b/ppp_comfyui.py @@ -215,6 +215,15 @@ class PromptPostProcessorComfyUINode: "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", + }, + ), "wc_options": ( "PPP_OPTIONS_WC", { @@ -310,6 +319,7 @@ class PromptPostProcessorComfyUINode: do_combinatorial, combinatorial_shuffle, combinatorial_limit, + combinatorial_randomsampler_fixed, model=None, wc_options=None, stn_options=None, @@ -426,6 +436,7 @@ class PromptPostProcessorComfyUINode: do_combinatorial=do_combinatorial, combinatorial_shuffle=combinatorial_shuffle, combinatorial_limit=combinatorial_limit, + combinatorial_randomsampler_fixed=combinatorial_randomsampler_fixed, results_file=results_file or "", ) self.wildcards_obj.refresh_wildcards( diff --git a/ppp_tree.py b/ppp_tree.py index 6c4f907..9855dc6 100644 --- a/ppp_tree.py +++ b/ppp_tree.py @@ -72,10 +72,9 @@ class TreeProcessor(lark.visitors.Interpreter): self.__insertion_at: list[tuple[int, int]] = [None for _ in range(10)] self.__detectedWildcards: list[tuple[str, bool]] = [] 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) @@ -127,37 +126,39 @@ class TreeProcessor(lark.visitors.Interpreter): self.__result = "" if not self.state.options.do_combinatorial: - self.__cycl_forced_path = list(self.state.cyclical_state.current_path) - self.__cycl_trace = [] + self.__forced_path = list(self.state.cyclical_state.current_path) + self.__trace = [] self.visit(parsed) self.__finalize_variables() - if self.__cycl_trace: - self.state.cyclical_state.last_trace = self.__cycl_trace[:] + 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())] # 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. + # __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, ...]: self.log(logging.DEBUG, f"Running combinatorial path: {forced_path}") - self.__comb_forced_path = list(forced_path) - self.__comb_trace = [] + 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())) 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) + return tuple(self.__trace) limit_reached = False @@ -1647,11 +1648,11 @@ class TreeProcessor(lark.visitors.Interpreter): 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) + decision_idx = len(self.__trace) + 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] @@ -1921,7 +1922,7 @@ class TreeProcessor(lark.visitors.Interpreter): weights = [weight + excluded_weights_sum / included_choices for weight in weights if weight >= 0] weights = np.array(weights) weights /= weights.sum() # normalize weights - comb_chosen_selection: Optional[list[dict]] = None + selected_choices: list[dict] = [] if available_choices: if from_value < 0: from_value = 1 @@ -1950,72 +1951,77 @@ class TreeProcessor(lark.visitors.Interpreter): all_selections.extend(combinations(available_choices, k)) else: all_selections.extend(permutations(available_choices, k)) - if specified_sampler == "~": - # It's a forced random sampler, so we randomly select one of the enumerated selections. - all_selections = [all_selections[self.__rng.choice(len(all_selections))]] - num_selections = len(all_selections) + decision_idx = len(self.__trace) if self.state.options.do_combinatorial: - decision_idx = len(self.__comb_trace) - self.__comb_trace.append(num_selections) + 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 + ): + # 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 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 num_choices < 2: + if chosen_num_choices < 2: repeating = False - comb_chosen_selection = ( - list(self.__rng.choice(available_choices, size=num_choices, p=weights, replace=repeating)) + 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 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: - selected_choices: list[dict] = comb_chosen_selection - 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 = [] + 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] container = options.get("container", None) if container is None: separator = options.get("separator", self.state.options.choice_separator) diff --git a/scripts/ppp_script.py b/scripts/ppp_script.py index 4cec2dc..329fa1c 100644 --- a/scripts/ppp_script.py +++ b/scripts/ppp_script.py @@ -181,6 +181,12 @@ class PromptPostProcessorA1111Script(scripts.Script): min_width=120, elem_id="ppp_combinatorial_limit", ) + combinatorial_randomsampler_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", + ) return [ force_equal_seeds, unlink_seed, @@ -189,6 +195,7 @@ class PromptPostProcessorA1111Script(scripts.Script): combinatorial, combinatorial_shuffle, combinatorial_limit, + combinatorial_randomsampler_fixed, ] def process( @@ -201,6 +208,7 @@ class PromptPostProcessorA1111Script(scripts.Script): input_combinatorial, input_combinatorial_shuffle, input_combinatorial_limit, + input_combinatorial_randomsampler_fixed, ): # pylint: disable=arguments-differ """ Processes the prompts and applies post-processing operations. @@ -214,6 +222,7 @@ class PromptPostProcessorA1111Script(scripts.Script): 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. Returns: None @@ -276,6 +285,7 @@ class PromptPostProcessorA1111Script(scripts.Script): 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), ) if not self.ppp_init: @@ -308,6 +318,7 @@ class PromptPostProcessorA1111Script(scripts.Script): "PPP prompt seed": input_seed, "PPP incremental seed": input_incremental_seed, "PPP combinatorial": input_combinatorial, + "PPP combinatorial random sampler fixed": input_combinatorial_randomsampler_fixed, } ) diff --git a/tests/base_tests.py b/tests/base_tests.py index f18e0b5..0760ef2 100644 --- a/tests/base_tests.py +++ b/tests/base_tests.py @@ -77,6 +77,7 @@ class TestPromptPostProcessorBase(unittest.TestCase): 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 "", ) self.def_env_info = { diff --git a/tests/tests_choices.py b/tests/tests_choices.py index bd6fbf3..cf16864 100644 --- a/tests/tests_choices.py +++ b/tests/tests_choices.py @@ -172,5 +172,30 @@ class TestChoices(TestPromptPostProcessorBase): OutputTuple("choice3, option1, a", ""), OutputTuple("choice3, option2, a", "", {"v": "option2"}), ], - combinatorial=True, + 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, + ), + ) + + def test_ch_combinatorial_random_consistent(self): # ~ sampler picks one value shared across all combinations + ppp_instance = self.init_ppp("nocup", combinatorial=True) + ppp_instance.process_prompts_group_start() + result = ppp_instance.process_prompt("{~a|b|c} {x|y}", "", 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}", )