* Improved model class detection.

* ComfyUI: ignore errors on install script.
This commit is contained in:
Antonio Cordero Balcazar
2026-06-08 13:07:20 +02:00
parent 6802f328f6
commit 423faf1336
5 changed files with 81 additions and 24 deletions
+5
View File
@@ -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
+10 -2
View File
@@ -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:
+11 -1
View File
@@ -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.
+12 -4
View File
@@ -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(
+43 -17
View File
@@ -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] = {}