* Additional model support.
* Support script to list models supported by host UIs.
This commit is contained in:
@@ -0,0 +1,314 @@
|
||||
"""
|
||||
Compares the model classes listed in the host's supported models file against
|
||||
the model definitions in ppp_config.yaml.defaults for a given host.
|
||||
|
||||
The relative path to the supported models file for each host is read from the
|
||||
comments preceding the `models:` key in the defaults file.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import ast
|
||||
import re
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import yaml
|
||||
|
||||
DEFAULT_CONFIG = Path(__file__).parent.parent / "ppp_config.yaml.defaults"
|
||||
SUPPORTED_HOSTS = {"comfyui", "reforge", "forge", "forgeneo", "sdnext"}
|
||||
|
||||
|
||||
def parse_host_file_paths(config_text: str) -> dict[str, str]:
|
||||
"""Extract the host->relative-path mapping from the comments above `models:`."""
|
||||
paths: dict[str, str] = {}
|
||||
in_block = False
|
||||
|
||||
for line in config_text.splitlines():
|
||||
stripped = line.strip()
|
||||
|
||||
if re.match(r"^#\s*Check supported models for each host in:", stripped):
|
||||
in_block = True
|
||||
continue
|
||||
|
||||
if in_block:
|
||||
# Lines like: "# host: some/relative/path.py" or "# host:"
|
||||
m = re.match(r"^#\s*(\w+)\s*:\s*(.*)", stripped)
|
||||
if m:
|
||||
host, rel_path = m.group(1), m.group(2).strip()
|
||||
if rel_path:
|
||||
paths[host] = rel_path
|
||||
else:
|
||||
# First non-matching line ends the block
|
||||
if not stripped.startswith("#"):
|
||||
break
|
||||
|
||||
return paths
|
||||
|
||||
|
||||
def extract_pipeline_classes(shared_items_path: Path) -> list[tuple[str, None]]:
|
||||
"""Parse shared_items.py and return unique diffusers pipeline class names from the pipelines dict."""
|
||||
source = shared_items_path.read_text(encoding="utf-8")
|
||||
tree = ast.parse(source, filename=str(shared_items_path))
|
||||
|
||||
class_names: set[str] = set()
|
||||
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.Assign):
|
||||
continue
|
||||
if not any(isinstance(t, ast.Name) and t.id == "pipelines" for t in node.targets):
|
||||
continue
|
||||
if not isinstance(node.value, ast.Dict):
|
||||
continue
|
||||
for value in node.value.values:
|
||||
# Match: getattr(diffusers, 'ClassName', None)
|
||||
if not isinstance(value, ast.Call):
|
||||
continue
|
||||
if not (isinstance(value.func, ast.Name) and value.func.id == "getattr"):
|
||||
continue
|
||||
if len(value.args) < 2:
|
||||
continue
|
||||
cls_arg = value.args[1]
|
||||
if isinstance(cls_arg, ast.Constant) and isinstance(cls_arg.value, str):
|
||||
class_names.add(cls_arg.value)
|
||||
|
||||
return [(name, None) for name in sorted(class_names)]
|
||||
|
||||
|
||||
def extract_model_classes(supported_models_path: Path) -> list[tuple[str, str | None]]:
|
||||
"""Parse a supported_models.py and return (class_name, parent_class) pairs in the `models` list(s)."""
|
||||
source = supported_models_path.read_text(encoding="utf-8")
|
||||
tree = ast.parse(source, filename=str(supported_models_path))
|
||||
|
||||
class_parents = _build_class_parents(tree)
|
||||
class_names: list[str] = []
|
||||
|
||||
for node in ast.walk(tree):
|
||||
# models = [ClassA, ClassB, ...]
|
||||
if isinstance(node, ast.Assign):
|
||||
for target in node.targets:
|
||||
if isinstance(target, ast.Name) and target.id == "models":
|
||||
class_names.extend(_names_from_list(node.value))
|
||||
|
||||
# models += [ClassA, ...]
|
||||
elif isinstance(node, ast.AugAssign):
|
||||
if isinstance(node.target, ast.Name) and node.target.id == "models":
|
||||
class_names.extend(_names_from_list(node.value))
|
||||
|
||||
sentinels = _find_sentinels(class_parents, set(class_names))
|
||||
|
||||
# Only keep classes that ultimately descend from a sentinel base.
|
||||
# Display parent is None when the immediate parent is a sentinel (class appears as a root).
|
||||
return [
|
||||
(name, None if class_parents.get(name) in sentinels else class_parents.get(name))
|
||||
for name in class_names
|
||||
if _has_base_ancestor(name, class_parents, sentinels)
|
||||
]
|
||||
|
||||
|
||||
def _build_class_parents(tree: ast.AST) -> dict[str, str | None]:
|
||||
"""Return a mapping of class name -> raw parent name."""
|
||||
parents: dict[str, str | None] = {}
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.ClassDef):
|
||||
continue
|
||||
parent: str | None = None
|
||||
for base in node.bases:
|
||||
if isinstance(base, ast.Attribute) and isinstance(base.value, ast.Name):
|
||||
parent = f"{base.value.id}.{base.attr}"
|
||||
elif isinstance(base, ast.Name):
|
||||
parent = base.id
|
||||
break # only the first base matters
|
||||
parents[node.name] = parent
|
||||
return parents
|
||||
|
||||
|
||||
def _find_sentinels(parents: dict[str, str | None], model_class_names: set[str]) -> set[str]:
|
||||
"""Detect root base class names: external refs (dotted) or local classes with no parent not in the models list."""
|
||||
sentinels: set[str] = set()
|
||||
for parent in parents.values():
|
||||
if parent is None:
|
||||
continue
|
||||
if "." in parent: # external reference, e.g. supported_models_base.BASE
|
||||
sentinels.add(parent)
|
||||
elif parent in parents and parents[parent] is None and parent not in model_class_names:
|
||||
sentinels.add(parent) # local class with no parent that is not itself a listed model
|
||||
return sentinels
|
||||
|
||||
|
||||
def _has_base_ancestor(name: str, parents: dict[str, str | None], sentinels: set[str]) -> bool:
|
||||
visited: set[str] = set()
|
||||
current = parents.get(name)
|
||||
while current is not None:
|
||||
if current in sentinels:
|
||||
return True
|
||||
if current in visited:
|
||||
return False # cycle guard
|
||||
visited.add(current)
|
||||
current = parents.get(current)
|
||||
return False
|
||||
|
||||
|
||||
def _names_from_list(node: ast.expr) -> list[str]:
|
||||
if not isinstance(node, ast.List):
|
||||
return []
|
||||
return [elt.id for elt in node.elts if isinstance(elt, ast.Name)]
|
||||
|
||||
|
||||
def _topo_sort_alpha(classes: list[tuple[str, str | None]]) -> list[tuple[str, str | None]]:
|
||||
"""Sort classes so each parent immediately precedes its children, with alphabetical ordering at every level."""
|
||||
class_set = {name for name, _ in classes}
|
||||
parent_of = {name: parent for name, parent in classes}
|
||||
|
||||
children_of: dict[str, list[str]] = {name: [] for name, _ in classes}
|
||||
roots: list[str] = []
|
||||
for name, parent in classes:
|
||||
if parent and parent in class_set:
|
||||
children_of[parent].append(name)
|
||||
else:
|
||||
roots.append(name)
|
||||
|
||||
roots.sort()
|
||||
for children in children_of.values():
|
||||
children.sort()
|
||||
|
||||
result: list[tuple[str, str | None]] = []
|
||||
|
||||
def visit(name: str) -> None:
|
||||
result.append((name, parent_of[name]))
|
||||
for child in children_of[name]:
|
||||
visit(child)
|
||||
|
||||
for root in roots:
|
||||
visit(root)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def build_class_to_model_map(config: dict, host: str) -> dict[str, str]:
|
||||
"""Return a mapping of class name -> model key for the given host."""
|
||||
mapping: dict[str, str] = {}
|
||||
models_config = config.get("models", {})
|
||||
|
||||
for model_key, model_data in models_config.items():
|
||||
if not isinstance(model_data, dict):
|
||||
continue
|
||||
detect = model_data.get("detect", {})
|
||||
if not isinstance(detect, dict):
|
||||
continue
|
||||
host_detect = detect.get(host)
|
||||
if not isinstance(host_detect, dict):
|
||||
continue
|
||||
classes = host_detect.get("class", [])
|
||||
if isinstance(classes, list):
|
||||
for cls in classes:
|
||||
mapping[cls] = model_key
|
||||
|
||||
return mapping
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Compare host model classes against ppp_config.yaml.defaults mappings."
|
||||
)
|
||||
parser.add_argument(
|
||||
"host",
|
||||
help="Host kind to compare against (e.g. comfyui, reforge, forge, ...).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"root",
|
||||
metavar="ROOT_FOLDER",
|
||||
type=Path,
|
||||
help="Root folder of the host UI installation.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--config",
|
||||
metavar="CONFIG_FILE",
|
||||
type=Path,
|
||||
default=DEFAULT_CONFIG,
|
||||
help=f"Path to ppp_config.yaml.defaults (default: {DEFAULT_CONFIG}).",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
root: Path = args.root
|
||||
host: str = args.host
|
||||
config_path: Path = args.config
|
||||
|
||||
if host not in SUPPORTED_HOSTS:
|
||||
parser.error(f"Host '{host}' is not supported by this script. Supported hosts: {', '.join(sorted(SUPPORTED_HOSTS))}")
|
||||
|
||||
if not root.is_dir():
|
||||
parser.error(f"ROOT_FOLDER does not exist or is not a directory: {root}")
|
||||
|
||||
if not config_path.exists():
|
||||
parser.error(f"Config file not found: {config_path}")
|
||||
|
||||
config_text = config_path.read_text(encoding="utf-8")
|
||||
host_paths = parse_host_file_paths(config_text)
|
||||
|
||||
known_hosts = list((yaml.safe_load(config_text).get("hosts") or {}).keys())
|
||||
if host not in known_hosts:
|
||||
print(f"Warning: '{host}' is not a recognized host. Known hosts: {', '.join(known_hosts)}")
|
||||
|
||||
rel_path = host_paths.get(host)
|
||||
if not rel_path:
|
||||
print(f"Error: no supported-models file path defined for host '{host}' in {config_path.name}.")
|
||||
sys.exit(1)
|
||||
|
||||
supported_models_path = root / rel_path.replace("\\", "/")
|
||||
if not supported_models_path.exists():
|
||||
print(f"Error: file not found: {supported_models_path}")
|
||||
sys.exit(1)
|
||||
|
||||
config = yaml.safe_load(config_text)
|
||||
if host == "sdnext":
|
||||
model_classes: list[tuple[str, str | None]] = extract_pipeline_classes(supported_models_path)
|
||||
else:
|
||||
model_classes = extract_model_classes(supported_models_path)
|
||||
if not model_classes:
|
||||
print("No model classes found in the models list.")
|
||||
sys.exit(1)
|
||||
|
||||
model_classes = _topo_sort_alpha(model_classes)
|
||||
|
||||
class_to_model = build_class_to_model_map(config, host)
|
||||
|
||||
class_set = {cls for cls, _ in model_classes}
|
||||
parent_of_display = {cls: parent for cls, parent in model_classes}
|
||||
depth_cache: dict[str, int] = {}
|
||||
|
||||
def get_depth(name: str) -> int:
|
||||
if name not in depth_cache:
|
||||
p = parent_of_display.get(name)
|
||||
depth_cache[name] = 0 if (not p or p not in class_set) else 1 + get_depth(p)
|
||||
return depth_cache[name]
|
||||
|
||||
# Build display labels: "ClassName" or "ClassName (Parent)", indented by depth
|
||||
labels = [f"{cls} ({parent})" if parent else cls for cls, parent in model_classes]
|
||||
depths = [get_depth(cls) for cls, _ in model_classes]
|
||||
col_width = max(d * 2 + len(lbl) for d, lbl in zip(depths, labels))
|
||||
missing: list[str] = []
|
||||
|
||||
print(f"Model classes in '{supported_models_path}' vs '{config_path.name}' (host: {host})")
|
||||
print("-" * (col_width + 42))
|
||||
|
||||
for (cls, _parent), label, depth in zip(model_classes, labels, depths):
|
||||
indented = " " * depth + label
|
||||
model = class_to_model.get(cls)
|
||||
if model:
|
||||
print(f"{indented:<{col_width}} -> {model}")
|
||||
else:
|
||||
missing.append((cls, _parent))
|
||||
print(f"{indented:<{col_width}} -> WARNING: not mapped")
|
||||
|
||||
print("-" * (col_width + 42))
|
||||
print(f"Total: {len(model_classes)} classes, {len(missing)} unmapped")
|
||||
|
||||
if missing:
|
||||
print("\nUnmapped classes:")
|
||||
for cls, parent in missing:
|
||||
label = f"{cls} ({parent})" if parent else cls
|
||||
print(f" - {label}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -304,13 +304,14 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
raise PPPInterrupt(errmsg)
|
||||
self.log(logging.WARNING, errmsg)
|
||||
|
||||
app = env_info.get("app", "")
|
||||
user_config_file = env_info.get("ppp_config", "")
|
||||
if isinstance(user_config_file, dict):
|
||||
user_cfg, _ = self.__parse_configuration(user_config_file, "forced configuration")
|
||||
else:
|
||||
user_raw: dict[str, Any] = {}
|
||||
if user_config_file == "":
|
||||
if env_info.get("app", "") == SUPPORTED_APPS.comfyui.value:
|
||||
if app == SUPPORTED_APPS.comfyui.value:
|
||||
try:
|
||||
import folder_paths # type: ignore
|
||||
|
||||
@@ -351,7 +352,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
self.known_models: list[str] = list(self.models_config.keys())
|
||||
|
||||
# Patch for tests (copy comfyui)
|
||||
if env_info.get("app", "") == "tests":
|
||||
if app == "tests":
|
||||
if self.config.hosts is None:
|
||||
self.config.hosts = {}
|
||||
self.config.hosts.setdefault("tests", HostConfig())
|
||||
@@ -362,10 +363,10 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
model.detect = {}
|
||||
model.detect.setdefault("tests", model.detect.get("comfyui", None))
|
||||
|
||||
host_config: HostConfig | None = (self.config.hosts or {}).get(env_info.get("app", ""))
|
||||
host_config: HostConfig | None = (self.config.hosts or {}).get(app)
|
||||
if host_config is None:
|
||||
raise PPPInterrupt(
|
||||
f"No host configuration found for app '{escape_single_quotes(env_info.get('app', ''))}'. Please check your configuration."
|
||||
f"No host configuration found for app '{escape_single_quotes(app)}'. Please check your configuration."
|
||||
)
|
||||
|
||||
# Update env_info with model detection
|
||||
@@ -381,7 +382,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
logging.WARNING,
|
||||
f"Variant name '{escape_single_quotes(v)}' in model '{escape_single_quotes(m)}' conflicts with a known model name. Discarding variant.",
|
||||
)
|
||||
self.log(logging.DEBUG, f"Host configuration: {host_config}", min_level=DEBUG_LEVEL.minimal)
|
||||
self.log(logging.DEBUG, f"Host configuration ({escape_single_quotes(app)}): {host_config}", min_level=DEBUG_LEVEL.minimal)
|
||||
|
||||
return host_config
|
||||
|
||||
@@ -390,7 +391,8 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
prop_base = env_info.get("property_base", None)
|
||||
model_class = env_info.get("model_class", "")
|
||||
model_name = env_info.get("model_filename", "")
|
||||
if not model_class and model_name and env_info.get("app", "") == SUPPORTED_APPS.comfyui.value:
|
||||
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)
|
||||
if model_class:
|
||||
env_info["model_class"] = model_class
|
||||
@@ -399,7 +401,6 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
|
||||
f"Detected model class '{model_class}' from filename '{model_name}'",
|
||||
min_level=DEBUG_LEVEL.minimal,
|
||||
)
|
||||
app = env_info.get("app", "")
|
||||
for m in self.known_models:
|
||||
env_info["is_" + m] = False
|
||||
model_obj = self.models_config.get(m)
|
||||
|
||||
+15
-15
@@ -79,12 +79,12 @@ hosts:
|
||||
|
||||
# Supported base models, variants, and options
|
||||
# Check supported models for each host in:
|
||||
# A1111:
|
||||
# Forge Classic: repositories\huggingface_guess\huggingface_guess\model_list.py
|
||||
# Forge Neo: modules_forge\packages\huggingface_guess\model_list.py
|
||||
# reForge: ldm_patched\modules\supported_models.py
|
||||
# SD.Next: modules\shared_items.py
|
||||
# ComfyUI: ComfyUI\comfy\supported_models.py
|
||||
# a1111:
|
||||
# forge: repositories\huggingface_guess\huggingface_guess\model_list.py
|
||||
# forgeneo: modules_forge\packages\huggingface_guess\model_list.py
|
||||
# reforge: ldm_patched\modules\supported_models.py
|
||||
# sdnext: modules\shared_items.py
|
||||
# comfyui: ComfyUI\comfy\supported_models.py
|
||||
models:
|
||||
# We define supported models and how we detect them in each host application (by class or by a known boolean property, or null for not supported).
|
||||
# We can also define here variants and some options.
|
||||
@@ -92,23 +92,23 @@ models:
|
||||
sd1: # Stable Diffusion 1
|
||||
detect:
|
||||
a1111: { property: "is_sd1" }
|
||||
forge: { property: "is_sd1" }
|
||||
forgeneo: { property: "is_sd1" }
|
||||
forge: { property: "is_sd1", class: ["SD15", "SD15_instructpix2pix"] }
|
||||
forgeneo: { property: "is_sd1", class: ["SD15"] }
|
||||
reforge: { property: "is_sd1" }
|
||||
sdnext: { class: ["LatentDiffusion", "StableDiffusionPipeline"] } # LatentDiffusion is for the original backend, StableDiffusionPipeline is for the diffusers backend; cannot differentiate SD1 and SD2, we set both to True
|
||||
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
|
||||
detect:
|
||||
a1111: { property: "is_sd2" }
|
||||
forge: { property: "is_sd2" }
|
||||
forge: { property: "is_sd2", class: ["SD20", "SD21UnclipL", "SD21UnclipH"] }
|
||||
forgeneo: null
|
||||
reforge: { property: "is_sd2" }
|
||||
sdnext: { class: ["LatentDiffusion", "StableDiffusionPipeline"] } # cannot differentiate SD1 and SD2, we set both to True; LatentDiffusion is for the original backend, StableDiffusionPipeline is for the diffusers backend
|
||||
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
|
||||
detect:
|
||||
a1111: { property: "is_ssd" }
|
||||
forge: null
|
||||
forge: { class: ["SSD1B"] }
|
||||
forgeneo: null
|
||||
reforge: { property: "is_ssd" }
|
||||
sdnext: null
|
||||
@@ -116,10 +116,10 @@ models:
|
||||
sdxl: # Stable Diffusion XL
|
||||
detect:
|
||||
a1111: { property: "is_sdxl" }
|
||||
forge: { property: "is_sdxl" }
|
||||
forgeneo: { 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" }
|
||||
sdnext: { class: ["StableDiffusionXLPipeline"] }
|
||||
sdnext: { class: ["StableDiffusionXLPipeline", "StableDiffusionXLImg2ImgPipeline", "StableDiffusionXLInpaintPipeline", "StableDiffusionXLInstructPix2PixPipeline"] }
|
||||
comfyui: { class: ["SDXL", "SDXLRefiner", "SDXL_instructpix2pix", "Segmind_Vega", "KOALA_700M", "KOALA_1B"] }
|
||||
variants:
|
||||
# At this level goes the name of the defined variants
|
||||
|
||||
Reference in New Issue
Block a user