Added API to classify checkpoint models
* SD1.5 / SD2.x / SDXL * inpaint model * refiner model
This commit is contained in:
@@ -1,2 +1,3 @@
|
||||
.vscode
|
||||
.env
|
||||
__pycache__
|
||||
@@ -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
@@ -1,4 +1,4 @@
|
||||
from . import nodes
|
||||
from . import api, nodes
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"ETN_LoadImageBase64": nodes.LoadImageBase64,
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user