diff --git a/README.md b/README.md index 0b589dc..e303960 100644 --- a/README.md +++ b/README.md @@ -163,7 +163,7 @@ There are various types of models that can be loaded as checkpoint, LoRA, Contro #### Paramters * `folder_name`: sub-directory in ComfyUI's models folder. - Supported model types: `checkpoints`, `diffusion_models` + Supported model types: `checkpoints`, `diffusion_models`, `unet`, `unet_gguf` #### Output Lists available models with additional classification info: @@ -177,11 +177,15 @@ Lists available models with additional classification info: ... } ``` -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, lumina2` +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, lumina2, chroma, qwen-image` 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. +Detection supports quantized models: +* GGUF: if the `gguf` module is installed, .gguf files are detected and will set the `quant` field +* Nunchaku: SVDQuant models are detected and will set the `quant` field to `svdq` + +Returns an entry `{"base_model": "unknown"}` for models with unknown format or which do not match any of the known base models. ### GET /api/etn/languages diff --git a/api.py b/api.py index 776097c..f99d34b 100644 --- a/api.py +++ b/api.py @@ -51,7 +51,8 @@ model_names = { "HiDream": "hi-dream", "Chroma": "chroma", "ACEStep": "ace-step", - "Omnigen2": "omnigen2" + "Omnigen2": "omnigen2", + "QwenImage": "qwen-image" } gguf_architectures = {"sd1": "sd15"} @@ -98,19 +99,32 @@ def inspect_safetensors(filename: str, model_type: str, is_checkpoint: bool): input_count = 4 # Find a matching base model depending on unet config - base_model = model_detection.model_config_from_unet_config(unet_config) + base_model = None + model_type = None + model_quant = None + if unet_config is not None: + base_model = model_detection.model_config_from_unet_config(unet_config) + if base_model is None: + raw_name = detect_svdq(cfg) + model_quant = "svdq" + else: + raw_name = base_model.__class__.__name__ + if raw_name == "SDXL": + model_type = base_model.model_type(cfg).name.lower().replace("_", "-") + + if not raw_name: return {"base_model": "unknown"} - base_model_class = base_model.__class__ - raw_name = base_model_class.__name__ base_model_name = model_names.get(raw_name, "unknown") result = {"base_model": base_model_name} result["is_inpaint"] = ( base_model_name in ["sd15", "sdxl"] and input_count > 4 ) or raw_name == "FluxInpaint" - if base_model_name == "sdxl": - result["type"] = base_model.model_type(cfg).name.lower().replace("_", "-") + if model_quant: + result["quant"] = model_quant + if model_type: + result["type"] = model_type elif "T2I" in raw_name: result["type"] = "t2i" elif "I2V" in raw_name: @@ -122,10 +136,24 @@ def inspect_safetensors(filename: str, model_type: str, is_checkpoint: bool): return result return {"base_model": "unknown"} except Exception as e: - # traceback.print_exc() + traceback.print_exc() return {"base_model": "unknown", "error": f"Failed to detect base model: {e}"} +def detect_svdq(cfg: dict) -> str | None: + if md := cfg.get("__metadata__"): + if comfy_config := md.get("comfy_config"): + if isinstance(comfy_config, str): + comfy_config = json.loads(comfy_config) + return comfy_config.get("model_class") + model_class = md.get("model_class") + if model_class == "NunchakuFluxTransformer2dModel": + return "Flux" + if model_class == "NunchakuQwenImageTransformer2DModel": + return "QwenImage" + return None + + def inspect_gguf(filename: str, model_type: str): try: import gguf @@ -146,10 +174,17 @@ def inspect_gguf(filename: str, model_type: str): return {"base_model": "flux", "is_inpaint": False} if arch_str == "flux" and any(t.name.startswith("distilled_guidance_layer") for t in itertools.islice(reader.tensors, 5)): arch_str = "chroma" - return { + + result = { "base_model": gguf_architectures.get(arch_str, arch_str), "is_inpaint": False, } + try: + result["quant"] = reader.get_field("general.file_type").lower() + except Exception as e: + result["quant"] = "gguf" + return result + except Exception as e: # traceback.print_exc() return {"base_model": "unknown", "error": f"Failed to detect base model: {e}"} diff --git a/pyproject.toml b/pyproject.toml index 51f3d14..0da672a 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 = "2.0.4" +version = "2.0.5" license = { file = "LICENSE" } [project.urls]