diff --git a/install.py b/install.py index 4b8a040..3836c61 100644 --- a/install.py +++ b/install.py @@ -1,3 +1,6 @@ +""" +Install dependencies for Prompt Post-Processor extension. For A1111 hosts. +""" from pathlib import Path requirements_filename = str(Path(__file__).resolve().parent / "requirements.txt") @@ -9,3 +12,5 @@ try: except ImportError: import launch launch.run_pip(f'install -r "{requirements_filename}"', "requirements for Prompt Post-Processor") +except Exception as e: + pass diff --git a/ppp.py b/ppp.py index c2094cf..a967651 100644 --- a/ppp.py +++ b/ppp.py @@ -20,6 +20,7 @@ from ppp_classes import ( HostConfig, ModelConfig, ModelDetectConfig, + PPPException, VariantConfig, PPPConfig, IFWILDCARDS_CHOICES, @@ -407,7 +408,14 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in model_name = env_info.get("model_filename", "") app = env_info.get("app", "") if not model_class and model_name and app == SUPPORTED_APPS.comfyui.value: - model_class = get_model_class_from_filename(model_name) + try: + model_class = get_model_class_from_filename(Path(model_name)) + except PPPException as e: + self.log( + logging.WARNING, + f"Could not detect model class from filename '{model_name}': {e}", + min_level=DEBUG_LEVEL.minimal, + ) if model_class: env_info["model_class"] = model_class self.log( @@ -1107,7 +1115,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in 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}") + self.log(logging.DEBUG, f"System variables: {filtered_sysvars}", DEBUG_LEVEL.minimal) self.log(logging.INFO, f"Combinatorial: {self.state.options.do_combinatorial}") def _expand_filename(self) -> Path: diff --git a/ppp_classes.py b/ppp_classes.py index 5577073..40da611 100644 --- a/ppp_classes.py +++ b/ppp_classes.py @@ -288,7 +288,17 @@ class PPPState: cyclical_state: CyclicalSamplerState = field(default_factory=CyclicalSamplerState) -class PPPInterrupt(Exception): +class PPPException(Exception): + """ + Custom exception to handle exceptions in the PromptPostProcessor. + """ + + def __init__(self, message: str = "An error occurred during prompt processing."): + super().__init__(message) + self.message = message + + +class PPPInterrupt(PPPException): """ Custom exception to handle interruptions in the PromptPostProcessor. This exception can be raised to stop the processing of prompts. diff --git a/ppp_comfyui.py b/ppp_comfyui.py index 897e508..cc8047a 100644 --- a/ppp_comfyui.py +++ b/ppp_comfyui.py @@ -11,7 +11,7 @@ import folder_paths # type: ignore import nodes # type: ignore from ppp import PromptPostProcessor -from ppp_classes import IFWILDCARDS_CHOICES, ONWARNING_CHOICES, SUPPORTED_APPS, PPPStateOptions +from ppp_classes import IFWILDCARDS_CHOICES, ONWARNING_CHOICES, SUPPORTED_APPS, PPPException, 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 @@ -323,8 +323,16 @@ class PromptPostProcessorComfyUINode: ) or "" if modelname == "(none)": modelname = "" - if modelclass == "": - modelclass = get_model_class_from_filename(modelname) + if modelclass == "" and modelname != "": + try: + modelclass = get_model_class_from_filename(Path(modelname)) + except PPPException as e: + log( + self.logger, + DEBUG_LEVEL.minimal, + logging.WARNING, + f"Could not detect model class from filename '{modelname}': {e}", + ) if modelclass: log( self.logger, @@ -337,7 +345,7 @@ class PromptPostProcessorComfyUINode: self.logger, DEBUG_LEVEL.minimal, logging.WARNING, - "Model class was not provided. System model variables will not be properly set.", + "Model class was not provided nor detected. System model variables will not be properly set.", ) if modelname == "": log( diff --git a/ppp_common.py b/ppp_common.py index 3639169..3f62222 100644 --- a/ppp_common.py +++ b/ppp_common.py @@ -10,7 +10,7 @@ import lark from ruamel.yaml import YAML as _YAML from ppp_logging import log -from ppp_classes import ONWARNING_CHOICES, PPPInterrupt, PPPState +from ppp_classes import ONWARNING_CHOICES, PPPException, PPPInterrupt, PPPState from ppp_utils import escape_single_quotes, format_output @@ -203,28 +203,42 @@ def preprocess_grammar(grammar_content: str, options: dict[str, bool], logger: l return "\n".join(result_lines) -def get_model_class_from_filename(filename: str) -> str: +def get_model_config_from_filename(filename: Path) -> object | None: + """ + Attempts to detect the model class from the given filename by inspecting the file header. + Currently only supports ComfyUI models in .safetensors format. + The path must be relative to a model folder. + """ try: import folder_paths # type: ignore import comfy.utils # type: ignore import comfy.model_detection as model_detection # type: ignore - except ImportError: - return "" + except ImportError as e: + raise PPPException(f"Error detecting class from '{filename}': {e}") from e import json if not filename: - return "" - full_path = ( - folder_paths.get_full_path("diffusion_models", filename) - or folder_paths.get_full_path("checkpoints", filename) - or folder_paths.get_full_path("unet", filename) - ) - if not full_path or not full_path.lower().endswith((".safetensors", ".sft")): - return "" + return None + + path_keys = ["diffusion_models", "checkpoints", "unet"] + full_path: Path | None = None + if filename.is_absolute(): + full_path = filename + else: + base_folders = [] + for key in path_keys: + base_folders.extend(folder_paths.get_folder_paths(key)) + for base in base_folders: + fname: Path = base / filename + if fname.exists(): + full_path = fname + break + if not full_path or not full_path.suffix.lower() in (".safetensors", ".sft"): + return None try: header_bytes = comfy.utils.safetensors_header(full_path) if header_bytes is None: - return "" + raise PPPException(f"Error detecting class from '{full_path}': no header") header = json.loads(header_bytes) # model_config_from_unet only inspects tensor shapes, not actual data. @@ -240,10 +254,22 @@ def get_model_class_from_filename(filename: str) -> str: mock_sd = {k: _ShapeProxy(v["shape"]) for k, v in header.items() if k != "__metadata__" and "shape" in v} prefix = model_detection.unet_prefix_from_state_dict(mock_sd) - config = model_detection.model_config_from_unet(mock_sd, prefix) - return config.__class__.__name__ if config else "" - except Exception: # pylint: disable=broad-except + 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") + return config + except Exception as e: # pylint: disable=broad-except + raise PPPException(f"Error detecting class from file '{full_path}': {e}") from e + + +def get_model_class_from_filename(filename: Path) -> str: + config = get_model_config_from_filename(filename) + if not config: return "" + c = config.__class__.__name__ + if not c: + raise PPPException(f"Error detecting class from '{filename}': config has no class name {config}") + return c def sanitize_wc_name(name: str) -> str: @@ -288,7 +314,7 @@ def convert_sdnext_styles_to_wildcard(inp: Path, out: Path): files = inp.glob("*.json") if inp.is_dir() else [inp] for file in files: with open(file, "r", encoding="utf-8-sig") as f: - data = _YAML(typ='safe').load(f) + data = _YAML(typ="safe").load(f) if not isinstance(data, list): continue wildcards[file] = {}