diff --git a/docs/COOKBOOK.md b/docs/COOKBOOK.md index c2066ff..ed2ae12 100644 --- a/docs/COOKBOOK.md +++ b/docs/COOKBOOK.md @@ -266,7 +266,7 @@ You can also specify a weight multiplier directly in the command - if both the c ## Send-to-negative with attention modifiers -When a send-to-negative command sits inside an attention modifier, the weight is carried over to the negative prompt. There is one important distinction: +When a send-to-negative command sits inside an attention modifier, the weight is carried over to the negative prompt. ```text (red appleround shape:1.4) @@ -276,17 +276,10 @@ Result in the negative: `(round shape:1.4)` - the surrounding weight is applied. ```text (red apple[round shape]:1.4) -``` - -Result in the negative: `([round shape]:1.4)` - the weight is **not** merged into the inner brackets, because the content of the negative tag is copied as-is and the merge happens before the tag is processed. - -```text (red [appleround shape]:1.4) ``` -Result in the negative: `(round shape:1.26)` - if merge attention is enabled, `1.4 × 0.9 = 1.26` is applied. - -Keep this in mind when writing choices in wildcards that combine attention brackets with send-to-negative. +For both cases, the result in the negative is: `(round shape:1.26)` - the weight is merged, if merge attention is enabled. ## Array variables diff --git a/docs/SYNTAX.md b/docs/SYNTAX.md index 5d44ec4..2789783 100644 --- a/docs/SYNTAX.md +++ b/docs/SYNTAX.md @@ -37,7 +37,7 @@ There is also a format where instead of `parameters$$` you just put the sampler, The construct parameters can be written with the following options (all are optional): -* "**~**" or "**@**": sampler (for compatibility with *Dynamic Prompts*), but only "**~**" (random) is supported. +* "**~**" (random) or "**@**" (cyclical): sampler (for compatibility with *Dynamic Prompts*). The cyclical sampler cycles through all combinations in order across consecutive `process_prompt` calls, resuming where the previous call left off (as long as the input prompt and negative prompt do not change). * "**r**": means it allows repetition of the choices. * "**o**": means it is "optional", and no error will be raised if there are no choices to select from. * "**n**" or "**n-m**" or "**n-**" or "**-m**": number or range of choices to select. Allows zero as the start of a range. Default is 1. @@ -471,7 +471,7 @@ They will be translated to the negative prompt. For example: * `(redsquare:1.5)` will end up as `(square:1.5)` in the negative prompt * `(red[square]:1.5)` will end up as `(square:1.35)` in the negative prompt (weight=1.5*0.9) if the merge attention option is enabled or `([square]:1.5)` otherwise. -* However `(red[square]:1.5)` will end up as `([square]:1.5)` in the negative prompt. The content of the negative tag is copied as is, and is not merged with the surrounding modifier because the insertions happen after the attention merging. +* `(red[square]:1.5)` will also end up as `(square:1.35)` in the negative prompt if the merge attention option is enabled, because the content of the negative tag is a single attention construct whose weight is merged with the surrounding modifier. ### Prompt editing constructs (alternation and scheduling) diff --git a/ppp.py b/ppp.py index f825e3f..73a166d 100644 --- a/ppp.py +++ b/ppp.py @@ -115,103 +115,8 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.logger = logger self.debug_level = options.debug_level self.interrupt_callback = interrupt - self.env_info = env_info - default_config_file = os.path.join(os.path.dirname(os.path.realpath(__file__)), "ppp_config.yaml.defaults") - try: - with open(default_config_file, "r", encoding="utf-8") as f: - default_raw: dict[str, Any] = yaml.safe_load(f) - except Exception as exc: # pylint: disable=broad-exception-caught - self.config = {} - raise PPPInterrupt( - f"Failed to load default configuration from '{escape_single_quotes(default_config_file)}'." - ) from exc - self.config, def_result = self.__parse_configuration(default_raw, "default configuration file") - if def_result != 0: - errmsg = "Default configuration file has errors. Please restore the default configuration file and, per instructions, use a copy to adapt it." - if def_result == 2: - raise PPPInterrupt(errmsg) - self.log(logging.WARNING, errmsg) - - user_config_file = self.env_info.get("ppp_config", "") - if isinstance(user_config_file, dict): - user_cfg, _ = self.__parse_configuration(user_config_file, "forced configuration") - else: - user_raw: dict[str, Any] = {} - if user_config_file == "": - if self.env_info.get("app", "") == SUPPORTED_APPS.comfyui.value: - try: - import folder_paths # type: ignore - - user_dir = folder_paths.get_user_directory() - if user_dir and os.path.isdir(user_dir): - user_config_file = os.path.join(user_dir, "default", "ppp_config.yaml") - except Exception: # pylint: disable=broad-exception-caught - self.log(logging.WARNING, "Failed to get user directory for PPP config.") - if not user_config_file or not os.path.exists(user_config_file): - user_config_file = os.path.join(os.path.dirname(os.path.realpath(__file__)), "ppp_config.yaml") - if user_config_file and os.path.exists(user_config_file): - with open(user_config_file, "r", encoding="utf-8") as f: - user_raw = yaml.safe_load(f) - user_cfg, _ = self.__parse_configuration(user_raw, "user configuration") - else: - user_cfg = None - if user_cfg is not None: - self.__merge_configuration(user_cfg) - - self.models_config: dict[str, ModelConfig | None] = self.config.models or {} - self.known_models: list[str] = list(self.models_config.keys()) - - # Patch for tests (copy comfyui) - if self.env_info.get("app", "") == "tests": - if self.config.hosts is None: - self.config.hosts = {} - self.config.hosts.setdefault("tests", HostConfig()) - for m in self.known_models: - model = self.models_config.get(m) - if model is not None: - if model.detect is None: - model.detect = {} - model.detect.setdefault("tests", model.detect.get("comfyui", None)) - - host_config: HostConfig | None = (self.config.hosts or {}).get(self.env_info.get("app", "")) - if host_config is None: - raise PPPInterrupt( - f"No host configuration found for app '{escape_single_quotes(self.env_info.get('app', ''))}'. Please check your configuration." - ) - - # Update env_info with model detection - prop_base = self.env_info.get("property_base", None) - model_class = self.env_info.get("model_class", "") - app = self.env_info.get("app", "") - for m in self.known_models: - self.env_info["is_" + m] = False - model_obj = self.models_config.get(m) - model_detect = (model_obj.detect if model_obj else None) or {} - model_detect_for_app: ModelDetectConfig | None = model_detect.get(app) - if model_detect_for_app is not None: - cls_list = model_detect_for_app.class_ or [] - if model_class in cls_list: - self.env_info["is_" + m] = True - elif model_detect_for_app.property is not None and prop_base is not None: - prop = model_detect_for_app.property - attr = getattr(prop_base, prop, None) - if isinstance(attr, bool) and attr: - self.env_info["is_" + m] = True - self.variants_definitions: dict[str, tuple[str, list[FindInFilenamePattern]]] = {} - for m in self.known_models: - model_obj = self.models_config.get(m) - for v, vo in ((model_obj.variants if model_obj else None) or {}).items(): - if v not in self.known_models: - self.variants_definitions[v] = (m, vo.find_in_filename) - else: - self.log( - logging.WARNING, - f"Variant name '{escape_single_quotes(v)}' in model '{escape_single_quotes(m)}' conflicts with a known model name. Discarding variant.", - ) - self.log(logging.DEBUG, f"Host configuration: {host_config}", min_level=DEBUG_LEVEL.minimal) - - # self.log(logging.INFO, f"Detected environment info: {env_info}", min_level=DEBUG_LEVEL.minimal) + host_config = self.__load_config_and_detect(env_info) if grammar_content is None: grammar_content = load_grammar() @@ -371,6 +276,128 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in def log(self, kind, message: str, min_level: DEBUG_LEVEL | None = None, exc_info: bool = False): log(self.logger, self.debug_level, kind, message, min_level, exc_info=exc_info) + def __load_config_and_detect(self, env_info: dict[str, Any]) -> HostConfig: + """Loads config files, performs model detection, and returns the resolved host config.""" + self.env_info = env_info + + default_config_file = os.path.join(os.path.dirname(os.path.realpath(__file__)), "ppp_config.yaml.defaults") + try: + with open(default_config_file, "r", encoding="utf-8") as f: + default_raw: dict[str, Any] = yaml.safe_load(f) + except Exception as exc: # pylint: disable=broad-exception-caught + self.config = {} + raise PPPInterrupt( + f"Failed to load default configuration from '{escape_single_quotes(default_config_file)}'." + ) from exc + self.config, def_result = self.__parse_configuration(default_raw, "default configuration file") + if def_result != 0: + errmsg = "Default configuration file has errors. Please restore the default configuration file and, per instructions, use a copy to adapt it." + if def_result == 2: + raise PPPInterrupt(errmsg) + self.log(logging.WARNING, errmsg) + + user_config_file = self.env_info.get("ppp_config", "") + if isinstance(user_config_file, dict): + user_cfg, _ = self.__parse_configuration(user_config_file, "forced configuration") + else: + user_raw: dict[str, Any] = {} + if user_config_file == "": + if self.env_info.get("app", "") == SUPPORTED_APPS.comfyui.value: + try: + import folder_paths # type: ignore + + user_dir = folder_paths.get_user_directory() + if user_dir and os.path.isdir(user_dir): + user_config_file = os.path.join(user_dir, "default", "ppp_config.yaml") + except Exception: # pylint: disable=broad-exception-caught + self.log(logging.WARNING, "Failed to get user directory for PPP config.") + if not user_config_file or not os.path.exists(user_config_file): + user_config_file = os.path.join(os.path.dirname(os.path.realpath(__file__)), "ppp_config.yaml") + if user_config_file and os.path.exists(user_config_file): + with open(user_config_file, "r", encoding="utf-8") as f: + user_raw = yaml.safe_load(f) + user_cfg, _ = self.__parse_configuration(user_raw, "user configuration") + else: + user_cfg = None + if user_cfg is not None: + self.__merge_configuration(user_cfg) + + self.models_config: dict[str, ModelConfig | None] = self.config.models or {} + self.known_models: list[str] = list(self.models_config.keys()) + + # Patch for tests (copy comfyui) + if self.env_info.get("app", "") == "tests": + if self.config.hosts is None: + self.config.hosts = {} + self.config.hosts.setdefault("tests", HostConfig()) + for m in self.known_models: + model = self.models_config.get(m) + if model is not None: + if model.detect is None: + model.detect = {} + model.detect.setdefault("tests", model.detect.get("comfyui", None)) + + host_config: HostConfig | None = (self.config.hosts or {}).get(self.env_info.get("app", "")) + if host_config is None: + raise PPPInterrupt( + f"No host configuration found for app '{escape_single_quotes(self.env_info.get('app', ''))}'. Please check your configuration." + ) + + # Update env_info with model detection + prop_base = self.env_info.get("property_base", None) + model_class = self.env_info.get("model_class", "") + app = self.env_info.get("app", "") + for m in self.known_models: + self.env_info["is_" + m] = False + model_obj = self.models_config.get(m) + model_detect = (model_obj.detect if model_obj else None) or {} + model_detect_for_app: ModelDetectConfig | None = model_detect.get(app) + if model_detect_for_app is not None: + cls_list = model_detect_for_app.class_ or [] + if model_class in cls_list: + self.env_info["is_" + m] = True + elif model_detect_for_app.property is not None and prop_base is not None: + prop = model_detect_for_app.property + attr = getattr(prop_base, prop, None) + if isinstance(attr, bool) and attr: + self.env_info["is_" + m] = True + self.variants_definitions: dict[str, tuple[str, list[FindInFilenamePattern]]] = {} + for m in self.known_models: + model_obj = self.models_config.get(m) + for v, vo in ((model_obj.variants if model_obj else None) or {}).items(): + if v not in self.known_models: + self.variants_definitions[v] = (m, vo.find_in_filename) + else: + self.log( + logging.WARNING, + f"Variant name '{escape_single_quotes(v)}' in model '{escape_single_quotes(m)}' conflicts with a known model name. Discarding variant.", + ) + self.log(logging.DEBUG, f"Host configuration: {host_config}", min_level=DEBUG_LEVEL.minimal) + + return host_config + + def update( + self, + env_info: dict[str, Any], + options: PPPStateOptions, + wildcards_obj: PPPWildcards, + extranetwork_mappings_obj: PPPExtraNetworkMappings, + ) -> None: + """Updates env_info, options, wildcards and enmappings while preserving the cyclical state and compiled parsers.""" + self.debug_level = options.debug_level + host_config = self.__load_config_and_detect(env_info) + self.state = PPPState( + logger=self.logger, + host_config=host_config, + options=options, + variables=VariableRepository(), + wildcards_obj=wildcards_obj, + extranetwork_mappings_obj=extranetwork_mappings_obj, + parsers=self.state.parsers, + cyclical_state=self.state.cyclical_state, + ) + self.__init_sysvars() + def __merge_configuration(self, user_config: PPPConfig): """ Merges the user configuration into the default configuration. @@ -968,6 +995,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.log(logging.INFO, f"Input negative_prompt: {negative_prompt}") self.log(logging.INFO, f"Combinatorial: {self.state.options.do_combinatorial}") t1 = time.monotonic_ns() + if self.state.cyclical_state.last_prompt_pair != (original_prompt, original_negative_prompt): + self.state.cyclical_state.reset() + self.state.cyclical_state.last_prompt_pair = (original_prompt, original_negative_prompt) results = self.__processprompts(np.random.default_rng(seed & 0xFFFFFFFF), prompt, negative_prompt) t2 = time.monotonic_ns() self.log(logging.INFO, f"Process prompt pair time: {(t2 - t1) / 1_000_000_000:.3f} seconds") diff --git a/ppp_classes.py b/ppp_classes.py index b4ed3db..7d60219 100644 --- a/ppp_classes.py +++ b/ppp_classes.py @@ -221,6 +221,37 @@ class PPPStateOptions: object.__setattr__(self, "cup_merge_attention", False) object.__setattr__(self, "cup_remove_extranetwork_tags", False) +class CyclicalSamplerState: + """Maintains the cycling position for '@' choice samplers across process_prompt calls.""" + + def __init__(self): + self.current_path: list[int] = [] + self.last_trace: list[int] = [] + self.last_prompt_pair: tuple[str, str] | None = None + + def advance(self): + """Advance to the next combination, cycling back to the start when all are exhausted.""" + if not self.last_trace: + self.current_path = [] + return + path = list(self.current_path) + while len(path) < len(self.last_trace): + path.append(0) + # Mixed-radix increment: least significant position is last. + for i in range(len(path) - 1, -1, -1): + path[i] += 1 + if path[i] < self.last_trace[i]: + break + path[i] = 0 + self.current_path = path + + def reset(self): + """Reset the cyclical state to the beginning.""" + self.current_path = [] + self.last_trace = [] + self.last_prompt_pair = None + + @dataclass(frozen=True) class PPPState: """State object passed to various PPP components during prompt processing.""" @@ -232,6 +263,7 @@ class PPPState: wildcards_obj: PPPWildcards = field(default_factory=PPPWildcards) extranetwork_mappings_obj: PPPExtraNetworkMappings = field(default_factory=PPPExtraNetworkMappings) parsers: dict[str, Lark] = field(default_factory=dict) + cyclical_state: CyclicalSamplerState = field(default_factory=CyclicalSamplerState) class PPPInterrupt(Exception): diff --git a/ppp_comfyui.py b/ppp_comfyui.py index d3aa4d9..62d5493 100644 --- a/ppp_comfyui.py +++ b/ppp_comfyui.py @@ -70,6 +70,7 @@ class PromptPostProcessorComfyUINode: self.grammar_content = file.read() self.wildcards_obj = PPPWildcards(lf.log) self.extranetwork_mappings_obj = PPPExtraNetworkMappings(lf.log) + self.ppp: PromptPostProcessor | None = None log( self.logger, DEBUG_LEVEL.minimal, @@ -394,16 +395,24 @@ class PromptPostProcessorComfyUINode: enmappings_folders, en_options["en_mappings_input"] if en_options else "", ) - ppp = PromptPostProcessor( - self.logger, - env_info, - options, - self.grammar_content, - self.interrupt, - self.wildcards_obj, - self.extranetwork_mappings_obj, - ) - results = ppp.process_prompt(pos_prompt, neg_prompt, seed if seed is not None else 1) + if self.ppp is None: + self.ppp = PromptPostProcessor( + self.logger, + env_info, + options, + self.grammar_content, + self.interrupt, + self.wildcards_obj, + self.extranetwork_mappings_obj, + ) + else: + self.ppp.update( + env_info, + options, + self.wildcards_obj, + self.extranetwork_mappings_obj, + ) + results = self.ppp.process_prompt(pos_prompt, neg_prompt, seed if seed is not None else 1) # with open(os.path.join(os.path.dirname(os.path.realpath(__file__)), "logs", "last_prompts_comfyui.txt"), "w", encoding="utf-8") as f: # f.write(f"Seed: {seed if seed is not None else 1}\n") diff --git a/ppp_tree.py b/ppp_tree.py index 0768333..1f17489 100644 --- a/ppp_tree.py +++ b/ppp_tree.py @@ -53,6 +53,8 @@ class TreeProcessor(lark.visitors.Interpreter): self.__result = "" self.__comb_forced_path: list[int] = [] self.__comb_trace: list[int] = [] + self.__cycl_forced_path: list[int] = [] + self.__cycl_trace: list[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) @@ -99,8 +101,13 @@ 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.visit(parsed) self.__finalize_echoed_variables() + if self.__cycl_trace: + self.state.cyclical_state.last_trace = self.__cycl_trace[:] + self.state.cyclical_state.advance() return [(self.__result, self.__detectedWildcards, self.state.variables.backup_user_and_echoed())] # Combinatorial mode: explore every possible path through choices and wildcards via DFS. @@ -150,9 +157,7 @@ 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 {limit} reached; some combinations have been skipped.") return results def __finalize_echoed_variables(self): @@ -1027,10 +1032,47 @@ class TreeProcessor(lark.visitors.Interpreter): parameters = str(negtagparameters) else: parameters = "" - content = self.__visit(tree.children[1::], False, True) + content_nodes = tree.children[1::] + attention_processing = self.state.host_config.attention + peeled = False + if ( + self.state.options.cup_merge_attention + and attention_processing in ("ok", "parentheses") + and len(content_nodes) == 1 + and isinstance(content_nodes[0], lark.Tree) + and content_nodes[0].data == "attention" + ): + inner_tree = content_nodes[0] + weight = 1.0 + while isinstance(inner_tree, lark.Tree) and inner_tree.data == "attention": + if len(inner_tree.children) == 2: + w = inner_tree.children[-1] + inner_weight = float(w) if w is not None else 1.1 + else: + inner_weight = 0.9 + weight *= inner_weight + inner_tree = inner_tree.children[0] + weight = math.floor(weight * 100) / 100 + weight_str = f"{weight:.2f}".rstrip("0").rstrip(".") + if weight_str == "0.9" and attention_processing != "parentheses": + weight_kind = 1 + elif weight_str == "1.1": + weight_kind = 2 + else: + weight_kind = 3 + if attention_processing == "parentheses" and weight_kind == 1: + weight_kind = 3 + weight_str = "0.9" + self.__shell.append(TreeProcessor.AccumulatedShell("at", (weight_kind, weight_str))) + content = self.__visit(inner_tree, False, True) + peeled = True + else: + content = self.__visit(content_nodes, False, True) self.__negtags.append( TreeProcessor.NegTag(len(self.__result), len(self.__result), content, parameters, self.__shell.copy()) ) + if peeled: + self.__shell.pop() info = f"with {escape_single_quotes(parameters) or 'no parameters'} : {escape_single_quotes(content)}" else: self.warn_or_stop("Ignored negative command in negative prompt") @@ -1624,7 +1666,7 @@ class TreeProcessor(lark.visitors.Interpreter): to_value: int = options.get("to", 1) separator: str = options.get("separator", self.state.options.choice_separator) msg_where = f"wildcard '{escape_single_quotes(wildcard_key)}'" if wildcard_key else "choices" - if sampler != "~": + if sampler not in ("~", "@"): self.warn_or_stop(f"Unsupported sampler '{escape_single_quotes(sampler)}' at {msg_where} options!") sampler = "~" expanded_choice_values = self.__get_choices_internal_get(choice_values, filter_specifier, wildcard_key) @@ -1659,7 +1701,7 @@ class TreeProcessor(lark.visitors.Interpreter): elif (to_value > len(available_choices) and not repeating) or from_value > to_value: to_value = len(available_choices) comb_chosen_selection: Optional[list[dict]] = None - if self.state.options.do_combinatorial: + if self.state.options.do_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, @@ -1679,13 +1721,22 @@ class TreeProcessor(lark.visitors.Interpreter): else: all_selections.extend(permutations(available_choices, k)) num_selections = len(all_selections) - decision_idx = len(self.__comb_trace) - self.__comb_trace.append(num_selections) - chosen_idx = ( - min(self.__comb_forced_path[decision_idx], num_selections - 1) - if decision_idx < len(self.__comb_forced_path) - else 0 - ) + if self.state.options.do_combinatorial: + decision_idx = len(self.__comb_trace) + self.__comb_trace.append(num_selections) + chosen_idx = ( + min(self.__comb_forced_path[decision_idx], num_selections - 1) + if decision_idx < len(self.__comb_forced_path) + else 0 + ) + else: # sampler == "@" + cycl_decision_idx = len(self.__cycl_trace) + self.__cycl_trace.append(num_selections) + chosen_idx = ( + self.__cycl_forced_path[cycl_decision_idx] % num_selections + if cycl_decision_idx < len(self.__cycl_forced_path) + else 0 + ) comb_chosen_selection = list(all_selections[chosen_idx]) num_choices = len(comb_chosen_selection) else: @@ -1705,7 +1756,7 @@ class TreeProcessor(lark.visitors.Interpreter): + (f" and separating with '{escape_single_quotes(separator)}'" if num_choices > 1 else ""), ) if num_choices > 0: - if self.state.options.do_combinatorial and comb_chosen_selection is not None: + if comb_chosen_selection is not None: selected_choices: list[dict] = comb_chosen_selection else: selected_choices: list[dict] = ( diff --git a/tests/tests_choices.py b/tests/tests_choices.py index b2ed022..ec6f26d 100644 --- a/tests/tests_choices.py +++ b/tests/tests_choices.py @@ -21,12 +21,62 @@ class TestChoices(TestPromptPostProcessorBase): ppp="nocup", ) - def test_ch_unsupportedsampler(self): # unsupported sampler + def test_ch_cyclical(self): # cyclical sampler cycles through all choices + ppp_instance = self.init_obj("nocup") self.process( InputTuple("the choices are: {@choice1|choice2|choice3}", ""), - OutputTuple("", ""), - ppp="nocup", - interrupted=True, + [ + OutputTuple("the choices are: choice1", ""), + OutputTuple("the choices are: choice2", ""), + OutputTuple("the choices are: choice3", ""), + OutputTuple("the choices are: choice1", ""), # cycles back + ], + ppp=ppp_instance, + ) + + def test_ch_cyclical_multiple_constructs(self): # two independent @ constructs cycle together + ppp_instance = self.init_obj("nocup") + self.process( + InputTuple("{@a|b} {@c|d}", ""), + [ + OutputTuple("a c", ""), + OutputTuple("a d", ""), + OutputTuple("b c", ""), + OutputTuple("b d", ""), + OutputTuple("a c", ""), # cycles back + ], + ppp=ppp_instance, + ) + + def test_ch_cyclical_resets_on_prompt_change(self): # state resets when the prompt pair changes + ppp_instance = self.init_obj("nocup") + # Advance the cycle to position 1 (choice2). + self.process( + InputTuple("the choices are: {@choice1|choice2|choice3}", ""), + [ + OutputTuple("the choices are: choice1", ""), + OutputTuple("the choices are: choice2", ""), + ], + ppp=ppp_instance, + ) + # A different prompt must restart from position 0 (choice1). + self.process( + InputTuple("the choices are: {@choice1|choice2|choice3} different", ""), + OutputTuple("the choices are: choice1 different", ""), + ppp=ppp_instance, + ) + + def test_ch_cyclical_mixed_samplers(self): # @ construct cycles while a ~ construct alongside is unaffected + ppp_instance = self.init_obj("nocup") + self.process( + InputTuple("{@a|b|c} {x|y}", ""), + [ + OutputTuple("a y", ""), + OutputTuple("b x", ""), + OutputTuple("c x", ""), + OutputTuple("a y", ""), # @ cycles back + ], + ppp=ppp_instance, ) def test_ch_choices_withcomments(self): # choices with comments and multiline diff --git a/tests/tests_stn.py b/tests/tests_stn.py index 50fe1a5..ddfb129 100644 --- a/tests/tests_stn.py +++ b/tests/tests_stn.py @@ -52,8 +52,7 @@ class TestSendToNegative(TestPromptPostProcessorBase): "normal quality", ), OutputTuple( - - "this is a ((test) (test:2):1.5) (red:1.5)", "[neg1], ([square]:1.5), normal quality, (neg2:1.65)" + "this is a ((test) (test:2):1.5) (red:1.5)", "[neg1], (square:1.35), normal quality, (neg2:1.65)" ), ) diff --git a/tests/tests_wildcards.py b/tests/tests_wildcards.py index 4fb373c..6d8f3a0 100644 --- a/tests/tests_wildcards.py +++ b/tests/tests_wildcards.py @@ -347,14 +347,6 @@ class TestWildcards(TestPromptPostProcessorBase): ppp="nocup", ) - def test_wc_unsupportedsampler(self): # unsupported sampler - self.process( - InputTuple("the choices are: __@yaml/wildcard2__", ""), - OutputTuple("", ""), - ppp="nocup", - interrupted=True, - ) - def test_wc_wildcard_globbing(self): # wildcard with globbing self.process( InputTuple("the choices are: __yaml/wildcard[12]__, __yaml/wildcard?__", ""), @@ -372,7 +364,7 @@ class TestWildcards(TestPromptPostProcessorBase): def test_wc_wildcardPS_yaml(self): # yaml wildcard with object formatted choices and options and prefix and suffix self.process( InputTuple("the choices are: __yaml/wildcardPS__", ""), - OutputTuple("the choices are: prefix-choice2/choice3-suffix", ""), + OutputTuple("the choices are: prefix1-choice2/choice3-suffix", ""), ppp="nocup", ) diff --git a/tests/wildcards/test.yaml b/tests/wildcards/test.yaml index f074320..65d664a 100644 --- a/tests/wildcards/test.yaml +++ b/tests/wildcards/test.yaml @@ -43,7 +43,7 @@ yaml: repeating: false, optional: false, count: 2, - prefix: "prefix-", + prefix: "prefix{1|2}-", suffix: "-suffix", separator: "/", }