From 8ed5591574ff6340729970eed780e79401568dad Mon Sep 17 00:00:00 2001 From: Acly Date: Sun, 12 Jan 2025 11:13:55 +0100 Subject: [PATCH] API breaking: removed is_refiner attribute from model inspection - sdxl refiner is reported with base model "sdxl-refiner" - added type attribute for sdxl model, allows to detect eps/v-prediction --- README.md | 6 ++++-- api.py | 19 ++++++++++--------- pyproject.toml | 2 +- 3 files changed, 15 insertions(+), 12 deletions(-) diff --git a/README.md b/README.md index de0e4cb..4e244d1 100644 --- a/README.md +++ b/README.md @@ -172,12 +172,14 @@ Lists available models with additional classification info: "checkpoint_file.safetensors": { "base_model": "sd15", "is_inpaint": false, - "is_refiner": false + "type": "eps" }, ... } ``` -Possible values for base model: `sd15, sd20, sd21, sd3, sdxl, ssd1b, svd, cascade-b, cascade-c, aura-flow, hunyuan-dit, flux, flux-schnell` +Possible values for base model: `sd15, sd20, sd21, sd3, sdxl, sdxl-refiner, ssd1b, svd, cascade-b, cascade-c, aura-flow, hunyuan-dit, flux, flux-schnell` + +If base model is `sdxl`, the `type` attribute is set with possible values: `eps, edm, v-prediction, v-prediction-edm` The entry is `{"base_model": "unknown"}` for models which are not in safetensors format or do not match any of the known base models. diff --git a/api.py b/api.py index d3b8c85..43eae84 100644 --- a/api.py +++ b/api.py @@ -22,7 +22,7 @@ model_names = { "SD20": "sd20", "SD21UnclipL": "sd21", "SD21UnclipH": "sd21", - "SDXLRefiner": "sdxl", + "SDXLRefiner": "sdxl-refiner", "SDXL": "sdxl", "SSD1B": "ssd1b", "SVD_img2vid": "svd", @@ -35,6 +35,9 @@ model_names = { "Flux": "flux", "FluxInpaint": "flux", "FluxSchnell": "flux-schnell", + "GenmoMochi": "mochi", + "LTXV": "ltxv", + "HunyuanVideo": "hunyuan-video", } gguf_architectures = {"sd1": "sd15"} @@ -89,14 +92,13 @@ def inspect_safetensors(filename: str, model_type: str, is_checkpoint: bool): base_model_class = base_model.__class__ base_model_name = model_names.get(base_model_class.__name__, "unknown") - is_inpaint = ( + result = {"base_model": base_model_name} + result["is_inpaint"] = ( base_model_name in ["sd15", "sdxl"] and input_count > 4 ) or base_model_class.__name__ == "FluxInpaint" - return { - "base_model": base_model_name, - "is_inpaint": is_inpaint, - "is_refiner": base_model_class is supported_models.SDXLRefiner, - } + if base_model_name == "sdxl": + result["type"] = base_model.model_type(cfg).name.lower().replace("_", "-") + return result return {"base_model": "unknown"} except Exception as e: # traceback.print_exc() @@ -120,11 +122,10 @@ def inspect_gguf(filename: str, model_type: str): ) arch_str = str(arch_field.parts[arch_field.data[-1]], encoding="utf-8") else: # stable-diffusion.cpp, requires conversion. not handled for now - return {"base_model": "flux", "is_inpaint": False, "is_refiner": False} + return {"base_model": "flux", "is_inpaint": False} return { "base_model": gguf_architectures.get(arch_str, arch_str), "is_inpaint": False, - "is_refiner": False, } except Exception as e: # traceback.print_exc() diff --git a/pyproject.toml b/pyproject.toml index d1498bf..adc5bea 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui-tooling-nodes" description = "Provides nodes and server API extensions geared towards using ComfyUI as a backend for external tools." -version = "1.6.0" +version = "2.0.0" license = { file = "LICENSE" } [project.urls]