diff --git a/dev/compare_models.py b/dev/compare_models.py new file mode 100644 index 0000000..1a8aded --- /dev/null +++ b/dev/compare_models.py @@ -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() diff --git a/ppp.py b/ppp.py index e86eaf4..a582a68 100644 --- a/ppp.py +++ b/ppp.py @@ -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) diff --git a/ppp_config.yaml.defaults b/ppp_config.yaml.defaults index 5f598a3..e702502 100644 --- a/ppp_config.yaml.defaults +++ b/ppp_config.yaml.defaults @@ -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