* 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:
@@ -9,3 +9,4 @@ tests/tests_local.py
|
||||
tests/local_wildcards
|
||||
tests/logs
|
||||
tools/*.bat
|
||||
dev/*.bat
|
||||
|
||||
@@ -85,11 +85,18 @@ See the [cookbook](docs/COOKBOOK.md) for interesting usages.
|
||||
|
||||
## Tools
|
||||
|
||||
A tool `tools/convert_styles.py` exists to convert A1111 or SD.Next styles into a wildcards file.
|
||||
### User tools
|
||||
|
||||
* `tools/convert_styles.py`: converts A1111 or SD.Next styles into a wildcards file.
|
||||
* `tools/check_loras.py`: checks for valid loras inside wildcards.
|
||||
|
||||
### Dev tools
|
||||
|
||||
* `dev/compare_models.py`: checks for new supported models in hosts.
|
||||
|
||||
## Contributing
|
||||
|
||||
To develop, I suggest doing so with the extension isolated from the UI (you can use a symlink to test it in the UI), and with its own virtual environment (venv or .venv), so the tests work and can be debugged properly.
|
||||
To develop, I suggest doing so with the extension isolated from the host UI (you can use a symlink to test it in the UI), and with its own virtual environment (venv or .venv), so the tests work and can be debugged properly.
|
||||
|
||||
## License
|
||||
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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):
|
||||
|
||||
@@ -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\]", ""),
|
||||
|
||||
Reference in New Issue
Block a user