Added API to classify checkpoint models

* SD1.5 / SD2.x / SDXL
* inpaint model
* refiner model
This commit is contained in:
Acly
2023-10-16 13:16:08 +02:00
parent c8c3667440
commit d7d5ee0ea8
4 changed files with 72 additions and 2 deletions
+1
View File
@@ -1,2 +1,3 @@
.vscode
.env
__pycache__
+24 -1
View File
@@ -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.
+1 -1
View File
@@ -1,4 +1,4 @@
from . import nodes
from . import api, nodes
NODE_CLASS_MAPPINGS = {
"ETN_LoadImageBase64": nodes.LoadImageBase64,
+46
View File
@@ -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)