Model inspection: support Qwen and Nunchaku quants
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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}"}
|
||||
|
||||
+1
-1
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user