* Improved model class detection.
* ComfyUI: ignore errors on install script.
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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] = {}
|
||||
|
||||
Reference in New Issue
Block a user