From d7d5ee0ea871fbda560168d9e015dd0ae2e13818 Mon Sep 17 00:00:00 2001 From: Acly Date: Mon, 16 Oct 2023 13:16:08 +0200 Subject: [PATCH] Added API to classify checkpoint models * SD1.5 / SD2.x / SDXL * inpaint model * refiner model --- .gitignore | 1 + README.md | 25 ++++++++++++++++++++++++- __init__.py | 2 +- api.py | 46 ++++++++++++++++++++++++++++++++++++++++++++++ 4 files changed, 72 insertions(+), 2 deletions(-) create mode 100644 api.py diff --git a/.gitignore b/.gitignore index 4edd750..6aa9490 100644 --- a/.gitignore +++ b/.gitignore @@ -1,2 +1,3 @@ .vscode +.env __pycache__ \ No newline at end of file diff --git a/README.md b/README.md index 62c2c78..4ea597b 100644 --- a/README.md +++ b/README.md @@ -1,6 +1,6 @@ # ComfyUI Nodes for External Tooling -Provides nodes geared towards using ComfyUI as a backend for external tools. +Provides nodes and API geared towards using ComfyUI as a backend for external tools. ## Nodes for sending and receiving images @@ -52,6 +52,29 @@ Copies a mask into the alpha channel of an image. * Inputs: image and mask * Outputs: RGBA image with mask used as transparency +## API for model inspection + +There are various types of models that can be loaded as checkpoint, LoRA, ControlNet, etc. which cannot be used interchangeably. The following API helps to categorize and filter them. + +### /etn/model_info + +Lists available models with additional classification info. +* Paramters: _none_ +* Output: list of model files + ``` + { + "checkpoint_file.safetensors": { + "base_model": "sd15"|"sd20"|"sdxl", + "is_inpaint": true|false, + "is_refiner": true|false + }, + ... + } + ``` + The entry is `{"base_model": "unknown"}` for models which are not in safetensors format. + +_Note: currently only supports checkpoints. May add other models in the future._ + ## Installation Download the repository and unpack into the `custom_nodes` folder in the ComfyUI installation directory. diff --git a/__init__.py b/__init__.py index 8e1a920..d2f95fc 100644 --- a/__init__.py +++ b/__init__.py @@ -1,4 +1,4 @@ -from . import nodes +from . import api, nodes NODE_CLASS_MAPPINGS = { "ETN_LoadImageBase64": nodes.LoadImageBase64, diff --git a/api.py b/api.py new file mode 100644 index 0000000..5c886eb --- /dev/null +++ b/api.py @@ -0,0 +1,46 @@ +import server +from aiohttp import web +import json + +import comfy.utils +import comfy.supported_models +import folder_paths +import server + +input_block = "model.diffusion_model.input_blocks.0.0.weight" +transformer_block = "1.transformer_blocks.0.attn2.to_k.weight" + + +def inspect_checkpoint(filename): + path = folder_paths.get_full_path("checkpoints", filename) + header = comfy.utils.safetensors_header(path) + if header: + cfg = json.loads(header.decode("utf-8")) + input_count = cfg[input_block]["shape"][1] + context_dim = next( + v["shape"][1] for k, v in cfg.items() if k.endswith(transformer_block) + ) + base_model = next( + model + for model in comfy.supported_models.models + if model.unet_config["context_dim"] == context_dim + ) + base_model_name = base_model.__name__.lower() + return { + "base_model": base_model_name[:4], + "is_inpaint": input_count > 4, + "is_refiner": "refiner" in base_model_name, + } + return {"base_model": "unknown"} + + +@server.PromptServer.instance.routes.get("/etn/model_info") +async def model_info(request): + try: + info = { + filename: inspect_checkpoint(filename) + for filename in folder_paths.get_filename_list("checkpoints") + } + return web.json_response(info) + except Exception as e: + return web.json_response(dict(error=str(e)), status=500)