* Improved detection of model class and some new models.

* Fixed seed size in ComfyUI to account for frontend limits.
* Better wildcard key validation.
* Some refactoring.
This commit is contained in:
Antonio Cordero Balcazar
2026-07-20 15:45:04 +02:00
parent 3a2571c008
commit 2eaf2b05fe
8 changed files with 147 additions and 29 deletions
+1
View File
@@ -9,3 +9,4 @@ tests/tests_local.py
tests/local_wildcards
tests/logs
tools/*.bat
dev/*.bat
+9 -2
View File
@@ -85,11 +85,18 @@ See the [cookbook](docs/COOKBOOK.md) for interesting usages.
## Tools
A tool `tools/convert_styles.py` exists to convert A1111 or SD.Next styles into a wildcards file.
### User tools
* `tools/convert_styles.py`: converts A1111 or SD.Next styles into a wildcards file.
* `tools/check_loras.py`: checks for valid loras inside wildcards.
### Dev tools
* `dev/compare_models.py`: checks for new supported models in hosts.
## Contributing
To develop, I suggest doing so with the extension isolated from the UI (you can use a symlink to test it in the UI), and with its own virtual environment (venv or .venv), so the tests work and can be debugged properly.
To develop, I suggest doing so with the extension isolated from the host UI (you can use a symlink to test it in the UI), and with its own virtual environment (venv or .venv), so the tests work and can be debugged properly.
## License
+14 -3
View File
@@ -34,7 +34,14 @@ from ppp_variables import VariableRepository, VariableEntry, VariableValue
from ppp_logging import DEBUG_LEVEL, log
from ppp_tree import TreeProcessor
from ppp_utils import escape_single_quotes, get_version_from_pyproject
from ppp_common import get_model_class_from_filename, load_grammar, parse_prompt, preprocess_grammar, warn_or_stop
from ppp_common import (
WARN_STOP_WHERE,
get_model_class_from_filename,
load_grammar,
parse_prompt,
preprocess_grammar,
warn_or_stop,
)
from ppp_wildcards import PPPWildcards
from ppp_enmappings import PPPExtraNetworkMappings
@@ -775,7 +782,11 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
self.log(logging.DEBUG, f"BREAK construct {break_replacements[break_processing][0]}")
elif break_processing == "error":
if re.search(r"\bBREAK\b", text):
warn_or_stop(self.state, where == -1, "BREAK constructs are not allowed!")
warn_or_stop(
self.state,
WARN_STOP_WHERE.negative if where == -1 else WARN_STOP_WHERE.positive,
"BREAK constructs are not allowed!",
)
if self.state.options.cup_ands:
# collapse ANDs with space after
@@ -955,7 +966,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
# Check for special character sequences that should not be in the result
compound_prompt = prompt + "\n" + negative_prompt
found_sequences = re.findall(r"::|\$\$|\$\{|[{}]", compound_prompt)
found_sequences = re.findall(r"::|\$\$|\$\{|[{}]|__", compound_prompt)
if found_sequences:
s = ", ".join(map(lambda x: '"' + x + '"', set(found_sequences)))
warnings.append(f"Probably invalid character sequences: {s}.")
+28 -4
View File
@@ -1,5 +1,6 @@
import ast
import csv
from enum import Enum
from functools import reduce
import logging
from pathlib import Path
@@ -82,13 +83,19 @@ def parse_prompt(
return parsed_prompt
def warn_or_stop(state: PPPState, is_negative: bool, message: str, e: Exception = None):
class WARN_STOP_WHERE(Enum):
none = 0
positive = 1
negative = 2
def warn_or_stop(state: PPPState, where: WARN_STOP_WHERE, message: str, e: Exception = None):
INVALID_CONTENT_STOP = "INVALID CONTENT! {0}\nBREAK "
if state.options.on_warning == ONWARNING_CHOICES.stop:
raise PPPInterrupt(
message,
INVALID_CONTENT_STOP.format(message) if not is_negative else "",
INVALID_CONTENT_STOP.format(message) if is_negative else "",
INVALID_CONTENT_STOP.format(message) if where == WARN_STOP_WHERE.positive else "",
INVALID_CONTENT_STOP.format(message) if where == WARN_STOP_WHERE.negative else "",
) from e
log(state.logger, state.options.debug_level, logging.WARNING, format_output(message))
@@ -256,8 +263,25 @@ def get_model_config_from_filename(filename: Path) -> object | None:
prefix = model_detection.unet_prefix_from_state_dict(mock_sd)
config = model_detection.model_config_from_unet(mock_sd, prefix, True)
if not config:
raise PPPException(f"Error detecting class from '{full_path}': no config found")
mock_sd, metadata = comfy.utils.convert_old_quants(mock_sd, "", metadata=None)
#Allow loading unets from checkpoint files
diffusion_model_prefix = model_detection.unet_prefix_from_state_dict(mock_sd)
temp_sd = comfy.utils.state_dict_prefix_replace(mock_sd, {diffusion_model_prefix: ""}, filter_keys=True)
if len(temp_sd) > 0:
mock_sd, metadata = comfy.utils.convert_old_quants(temp_sd, "", metadata=metadata)
config = model_detection.model_config_from_unet(mock_sd, "", metadata=metadata)
if config is None:
mock_sd = model_detection.convert_diffusers_mmdit(mock_sd, "")
if mock_sd is not None: #diffusers mmdit
config = model_detection.model_config_from_unet(mock_sd, "")
else: #diffusers unet
config = model_detection.model_config_from_diffusers_unet(mock_sd)
if not config:
raise PPPException(f"Error detecting class from file '{full_path}': no config found")
return config
except PPPException:
raise
except Exception as e: # pylint: disable=broad-except
raise PPPException(f"Error detecting class from file '{full_path}': {e}") from e
+39 -7
View File
@@ -82,7 +82,7 @@ hosts:
alternation: error
and: comma
break: comma
seed_bits: 64
seed_bits: 53 # and not 64, because of frontend JavaScript limits
# Supported base models, variants, and options
# Check supported models for each host in:
@@ -101,7 +101,7 @@ models:
a1111: { property: "is_sd1" }
forge: { property: "is_sd1", class: ["SD15", "SD15_instructpix2pix"] }
forgeneo: { property: "is_sd1", class: ["SD15"] }
reforge: { property: "is_sd1" }
reforge: { property: "is_sd1", class: ["SD15", "SD15_instructpix2pix"] }
sdnext: { class: ["LatentDiffusion", "StableDiffusionPipeline", "StableDiffusionInpaintPipeline", "StableDiffusionInstructPix2PixPipeline", "StableDiffusionUpscalePipeline"] } # LatentDiffusion is for the original backend, StableDiffusionPipeline is for the diffusers backend; cannot differentiate SD1 and SD2, we set both to True
comfyui: { class: ["SD15", "SD15_instructpix2pix"] }
sd2: # Stable Diffusion 2
@@ -109,7 +109,7 @@ models:
a1111: { property: "is_sd2" }
forge: { property: "is_sd2", class: ["SD20", "SD21UnclipL", "SD21UnclipH"] }
forgeneo: null
reforge: { property: "is_sd2" }
reforge: { property: "is_sd2", class: ["SD20", "SD21UnclipL", "SD21UnclipH"] }
sdnext: { class: ["LatentDiffusion", "StableDiffusionPipeline", "StableDiffusionInpaintPipeline", "StableDiffusionInstructPix2PixPipeline", "StableDiffusionUpscalePipeline"] } # cannot differentiate SD1 and SD2, we set both to True; LatentDiffusion is for the original backend, StableDiffusionPipeline is for the diffusers backend
comfyui: { class: ["SD20", "SD21UnclipL", "SD21UnclipH", "LotusD"] }
ssd: # Segmind Stable Diffusion 1B
@@ -117,7 +117,7 @@ models:
a1111: { property: "is_ssd" }
forge: { class: ["SSD1B"] }
forgeneo: null
reforge: { property: "is_ssd" }
reforge: { property: "is_ssd", class: ["SSD1B"] }
sdnext: null
comfyui: { class: ["SSD1B"]}
sdxl: # Stable Diffusion XL
@@ -125,7 +125,7 @@ models:
a1111: { property: "is_sdxl" }
forge: { property: "is_sdxl", class: ["SDXL", "SDXLRefiner", "SDXL_instructpix2pix", "Segmind_Vega", "KOALA_700M", "KOALA_1B"] }
forgeneo: { property: "is_sdxl", class: ["SDXL", "SDXLRefiner"] }
reforge: { property: "is_sdxl" }
reforge: { property: "is_sdxl", class: ["SDXL", "SDXLRefiner", "SDXL_instructpix2pix", "Segmind_Vega", "KOALA_700M", "KOALA_1B"] }
sdnext: { class: ["StableDiffusionXLPipeline", "StableDiffusionXLImg2ImgPipeline", "StableDiffusionXLInpaintPipeline", "StableDiffusionXLInstructPix2PixPipeline"] }
comfyui: { class: ["SDXL", "SDXLRefiner", "SDXL_instructpix2pix", "Segmind_Vega", "KOALA_700M", "KOALA_1B"] }
variants:
@@ -142,7 +142,7 @@ models:
a1111: { property: "is_sd3" }
forge: { property: "is_sd3", class: ["SD3"] }
forgeneo: null
reforge: { property: "is_sd3" }
reforge: { property: "is_sd3", class: ["SD3"] }
sdnext: { class: ["StableDiffusion3Pipeline"] }
comfyui: { class: ["SD3"] }
flux: # Flux 1
@@ -248,7 +248,7 @@ models:
forgeneo: { class: ["WAN21_T2V", "WAN21_I2V"] }
reforge: { class: ["WAN21_T2V", "WAN21_I2V", "WAN21_FunControl2V", "WAN21_Camera", "WAN22_Camera", "WAN21_Vace", "WAN21_HuMo", "WAN22_S2V", "WAN22_Animate", "WAN22_T2V"] }
sdnext: { class: ["WanPipeline"] }
comfyui: { class: ["WAN21_T2V", "WAN21_I2V", "WAN21_FunControl2V", "WAN21_Camera", "WAN22_Camera", "WAN21_Vace", "WAN21_HuMo", "WAN22_S2V", "WAN22_Animate", "WAN22_T2V", "WAN21_FlowRVS", "WAN21_SCAIL", "WAN22_WanDancer"] }
comfyui: { class: ["WAN21_T2V", "WAN21_I2V", "WAN21_FunControl2V", "WAN21_Camera", "WAN22_Camera", "WAN21_Vace", "WAN21_HuMo", "WAN22_S2V", "WAN22_Animate", "WAN22_T2V", "WAN21_FlowRVS", "WAN21_SCAIL", "WAN22_WanDancer", "WAN21_CausalAR_T2V", "WAN21_SCAIL2"] }
hidream: # HiDream
detect:
a1111: null
@@ -337,3 +337,35 @@ models:
reforge: null
sdnext: null
comfyui: { class: ["CogVideoX_T2V", "CogVideoX_I2V", "CogVideoX_Inpaint"] }
krea2: # Krea2
detect:
a1111: null
forge: null
forgeneo: null
reforge: null
sdnext: null
comfyui: { class: ["Krea2"] }
Ideogram4: # Ideogram4
detect:
a1111: null
forge: null
forgeneo: null
reforge: null
sdnext: null
comfyui: { class: ["Ideogram4"] }
Lens: # Lens
detect:
a1111: null
forge: null
forgeneo: null
reforge: null
sdnext: null
comfyui: { class: ["Lens"] }
Boogu: # Boogu
detect:
a1111: null
forge: null
forgeneo: null
reforge: null
sdnext: null
comfyui: { class: ["Boogu"] }
+7 -2
View File
@@ -15,7 +15,7 @@ from ppp_classes import IFWILDCARDS_CHOICES, SUPPORTED_APPS, PPPState
from ppp_enmappings import PPPENMappingVariant
from ppp_logging import DEBUG_LEVEL, log
from ppp_utils import escape_single_quotes, repr_value
from ppp_common import parse_prompt, warn_or_stop
from ppp_common import WARN_STOP_WHERE, parse_prompt, warn_or_stop
from ppp_variables import ScalarValue, VariableEntry
from ppp_wildcards import PPPWildcard
@@ -81,7 +81,12 @@ class TreeProcessor(lark.visitors.Interpreter):
log(self.state.logger, self.state.options.debug_level, kind, message, min_level)
def warn_or_stop(self, message: str, e: Exception = None):
warn_or_stop(self.state, self.__is_negative, message, e)
warn_or_stop(
self.state,
WARN_STOP_WHERE.negative if self.__is_negative else WARN_STOP_WHERE.positive,
message,
e,
)
def __reset_run_state(self):
"""Reset all per-run mutable state for a fresh combinatorial pass."""
+45 -8
View File
@@ -1,5 +1,6 @@
import fnmatch
from pathlib import Path
import re
from typing import Any, Optional
import logging
from ruamel.yaml import YAML as _YAML
@@ -114,23 +115,33 @@ class PPPWildcards:
keys = sorted(fnmatch.filter(self.wildcards.keys(), key))
return [self.wildcards[k] for k in keys]
def __get_wc_in_dict(self, dictionary: dict, prefix="") -> list[tuple[str, Any]]:
def __get_wc_in_dict(self, dictionary: dict, prefix="", file_str: str = "") -> list[tuple[str, Any]]:
"""
Get all wildcards in a dictionary, along their object.
Args:
dictionary (dict): The dictionary to check.
prefix (str): The prefix for the current key.
file_str (str): The file string for logging purposes.
Returns:
list: A list of all leaf wildcards in the dictionary.
"""
wc = []
for key, obj in dictionary.items():
strkey = str(key)
if not self.__check_key_validity(strkey):
log(
self.__logger,
self.__debug_level,
logging.WARNING,
f"Invalid wildcard name part '{escape_single_quotes(prefix + strkey)}' in file '{escape_single_quotes(file_str)}'!",
)
continue
if isinstance(obj, dict):
wc.extend(self.__get_wc_in_dict(obj, prefix + str(key) + "/"))
wc.extend(self.__get_wc_in_dict(obj, prefix + strkey + "/", file_str))
else:
wc.append((prefix + str(key), obj))
wc.append((prefix + strkey, obj))
return wc
def __remove_wildcards_from_path(self, full_path: Path, debug=True):
@@ -218,7 +229,7 @@ class PPPWildcards:
wildcards_input = wildcards_input.strip()
if wildcards_input != "":
try:
content = _YAML(typ='safe').load(wildcards_input)
content = _YAML(typ="safe").load(wildcards_input)
except _YAMLError as e:
log(self.__logger, self.__debug_level, logging.WARNING, f"Invalid format for input wildcards: {e}")
return
@@ -374,6 +385,21 @@ class PPPWildcards:
value = f"{options}::{value}"
return value
def __check_key_validity(self, key: str) -> bool:
"""
Check if a key is valid.
Args:
key (str): The key to check.
Returns:
bool: Whether the key is valid or not.
"""
match = re.match(r"^[a-zA-Z0-9\-_#]+$", key)
if match is None:
return False
return True
def __add_wildcard(self, content: object, full_path: Path | None, external_key_parts: list[str]):
"""
Add a wildcard to the wildcards dictionary.
@@ -389,9 +415,20 @@ class PPPWildcards:
return str(wc.file) if wc.file is not None else "input"
key_parts = external_key_parts.copy()
if key_parts:
for part in key_parts:
strkey = str(part)
if strkey and not self.__check_key_validity(strkey):
log(
self.__logger,
self.__debug_level,
logging.WARNING,
f"Invalid wildcard name start '{escape_single_quotes('/'.join(key_parts))}' in file '{escape_single_quotes(file_str)}'!",
)
return
if isinstance(content, dict):
key_parts.pop()
keys = self.__get_wc_in_dict(content)
key_parts.pop() # we don't want the name of the filename to be included, since the dict keys will be used instead
keys = self.__get_wc_in_dict(content, "", file_str)
for key, obj in keys:
tmp_key_parts = key_parts.copy()
tmp_key_parts.extend(key.split("/"))
@@ -472,7 +509,7 @@ class PPPWildcards:
external_key_parts = list(full_path.with_suffix("").relative_to(base).parts)
try:
with open(full_path, "r", encoding="utf-8") as file:
content = _YAML(typ='safe').load(file)
content = _YAML(typ="safe").load(file)
except: # pylint: disable=bare-except
log(
self.__logger,
@@ -481,7 +518,7 @@ class PPPWildcards:
f"Could not read file '{escape_single_quotes(str(full_path))}' with utf-8 encoding, trying windows-1252...",
)
with open(full_path, "r", encoding="windows-1252") as file:
content = _YAML(typ='safe').load(file)
content = _YAML(typ="safe").load(file)
self.__add_wildcard(content, full_path, external_key_parts)
def __get_wildcards_in_text_file(self, full_path: Path, base: Path):
+4 -3
View File
@@ -1,10 +1,9 @@
import logging
from dataclasses import replace
from ppp import PromptPostProcessor # type: ignore
from ppp import PromptPostProcessor
from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase
if __name__ == "__main__":
raise SystemExit("This script must not be run directly")
@@ -191,7 +190,9 @@ class TestCleanup(TestPromptPostProcessorBase):
"Expected an 'Unmatched' warning",
)
def test_cl_warn_escaped_unmatched_no_false_warning(self): # escaped unmatched paren/bracket does not trigger warning
def test_cl_warn_escaped_unmatched_no_false_warning(
self,
): # escaped unmatched paren/bracket does not trigger warning
with self.assertNoLogs("PromptPostProcessor", level=logging.WARNING):
self.process(
InputTuple(r"text with \(escaped unmatched\]", ""),