mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-18 19:41:26 -03:00
fc3f3f3bdb
The Checkpoint/Unet Loader (LoraManager) nodes now support ComfyUI's built-in control_after_generate mechanism on the ckpt_name/unet_name combos, letting users pick a random model on every queue with the selected model written back into the widget (visible, and lockable via the 'fixed' mode). A base_model input narrows the random pool: a front-end extension fetches the name/base_model mapping from the new /api/lm/checkpoints/loader-pool endpoint and filters the combo options, wired through the node callback, the refreshComboInNodes extension hook, and a graph.onConfigure hook installed from onAdded (onNodeCreated fires before the node is attached to a graph, so the graph reference is unavailable there).
305 lines
12 KiB
Python
305 lines
12 KiB
Python
import logging
|
|
import os
|
|
from typing import Any, List, Tuple
|
|
import comfy.sd # pyright: ignore[reportMissingImports]
|
|
from ..utils.utils import get_checkpoint_info_absolute, _format_model_name_for_comfyui
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def _reload_gguf_unet(
|
|
unet_path: str, weight_dtype: str, disable_dynamic: bool = False
|
|
) -> object:
|
|
"""Reload a GGUF diffusion model from disk (cached_patcher_init factory).
|
|
|
|
Mirrors the GGUF branch of UNETLoaderLM.load_unet so ModelPatcher
|
|
deepclone/dynamic machinery can rebuild GGUF models with the correct
|
|
GGMLOps. ``disable_dynamic`` is accepted for signature compatibility
|
|
with core ComfyUI loaders.
|
|
"""
|
|
loader = UNETLoaderLM()
|
|
model, = loader._load_gguf_unet(unet_path, unet_path, weight_dtype)
|
|
return model
|
|
|
|
|
|
class UNETLoaderLM:
|
|
"""UNET Loader with support for extra folder paths
|
|
|
|
Loads diffusion models/UNets from both standard ComfyUI folders and LoRA Manager's
|
|
extra folder paths, providing a unified interface for UNET loading.
|
|
Supports both regular diffusion models and GGUF format models.
|
|
The unet_name combo supports ComfyUI's control_after_generate, letting
|
|
users pick a random diffusion model on every run; the base_model input
|
|
narrows the random pool through a front-end extension that filters the
|
|
combo options.
|
|
"""
|
|
|
|
NAME = "Unet Loader (LoraManager)"
|
|
CATEGORY = "Lora Manager/loaders"
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
# Get list of unet names from scanner (includes extra folder paths)
|
|
unet_names = cls._get_unet_names()
|
|
base_models = cls._get_available_base_models()
|
|
return {
|
|
"required": {
|
|
"unet_name": (
|
|
unet_names,
|
|
{
|
|
"tooltip": (
|
|
"The name of the diffusion model to load. Use "
|
|
"control_after_generate to pick a random model on "
|
|
"every run."
|
|
),
|
|
"control_after_generate": True,
|
|
},
|
|
),
|
|
"weight_dtype": (
|
|
["default", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e5m2"],
|
|
{"tooltip": "The dtype to use for the model weights."},
|
|
),
|
|
"base_model": (
|
|
base_models,
|
|
{
|
|
"default": "Any",
|
|
"tooltip": (
|
|
"Restrict the random selection pool to this base "
|
|
"model. 'Any' uses the full pool."
|
|
),
|
|
},
|
|
),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("MODEL",)
|
|
RETURN_NAMES = ("MODEL",)
|
|
OUTPUT_TOOLTIPS = ("The model used for denoising latents.",)
|
|
FUNCTION = "load_unet"
|
|
|
|
@classmethod
|
|
def _get_unet_names(cls) -> List[str]:
|
|
"""Get list of diffusion model names from scanner cache in ComfyUI format (relative path with extension)"""
|
|
try:
|
|
from ..services.service_registry import ServiceRegistry
|
|
import asyncio
|
|
|
|
async def _get_names():
|
|
scanner = await ServiceRegistry.get_checkpoint_scanner()
|
|
cache = await scanner.get_cached_data()
|
|
|
|
# Get all model roots for calculating relative paths
|
|
model_roots = scanner.get_model_roots()
|
|
|
|
# Filter only diffusion_model type and format names
|
|
names = []
|
|
for item in cache.raw_data:
|
|
if item.get("sub_type") == "diffusion_model":
|
|
file_path = item.get("file_path", "")
|
|
# Only offer models that still exist on disk so ComfyUI
|
|
# flags missing diffusion models at queue time via
|
|
# "value not in list" (the scanner cache can be stale).
|
|
if file_path and os.path.exists(file_path):
|
|
# Format using relative path with OS-native separator
|
|
formatted_name = _format_model_name_for_comfyui(
|
|
file_path, model_roots
|
|
)
|
|
if formatted_name:
|
|
names.append(formatted_name)
|
|
|
|
return sorted(names)
|
|
|
|
try:
|
|
loop = asyncio.get_running_loop()
|
|
import concurrent.futures
|
|
|
|
def run_in_thread():
|
|
new_loop = asyncio.new_event_loop()
|
|
asyncio.set_event_loop(new_loop)
|
|
try:
|
|
return new_loop.run_until_complete(_get_names())
|
|
finally:
|
|
new_loop.close()
|
|
|
|
with concurrent.futures.ThreadPoolExecutor() as executor:
|
|
future = executor.submit(run_in_thread)
|
|
return future.result()
|
|
except RuntimeError:
|
|
return asyncio.run(_get_names())
|
|
except Exception as e:
|
|
logger.error(f"Error getting unet names: {e}")
|
|
return []
|
|
|
|
@classmethod
|
|
def _get_available_base_models(cls) -> List[str]:
|
|
"""Get distinct base_model values present among indexed diffusion models, for the random-selection filter."""
|
|
try:
|
|
from ..services.service_registry import ServiceRegistry
|
|
|
|
async def _get_base_models():
|
|
scanner = await ServiceRegistry.get_checkpoint_scanner()
|
|
cache = await scanner.get_cached_data()
|
|
|
|
base_models = set()
|
|
for item in cache.raw_data:
|
|
if item.get("sub_type") != "diffusion_model":
|
|
continue
|
|
base_model = item.get("base_model")
|
|
file_path = item.get("file_path", "")
|
|
if base_model and file_path and os.path.exists(file_path):
|
|
base_models.add(base_model)
|
|
|
|
return sorted(base_models)
|
|
|
|
return ["Any"] + cls._run_async(_get_base_models)
|
|
except Exception as e:
|
|
logger.error(f"Error getting available base models: {e}")
|
|
return ["Any"]
|
|
|
|
@staticmethod
|
|
def _run_async(coro_fn):
|
|
"""Run an async fetcher, handling the case where an event loop is already running."""
|
|
import asyncio
|
|
|
|
try:
|
|
asyncio.get_running_loop()
|
|
import concurrent.futures
|
|
|
|
def run_in_thread():
|
|
new_loop = asyncio.new_event_loop()
|
|
asyncio.set_event_loop(new_loop)
|
|
try:
|
|
return new_loop.run_until_complete(coro_fn())
|
|
finally:
|
|
new_loop.close()
|
|
|
|
with concurrent.futures.ThreadPoolExecutor() as executor:
|
|
future = executor.submit(run_in_thread)
|
|
return future.result()
|
|
except RuntimeError:
|
|
return asyncio.run(coro_fn())
|
|
|
|
def load_unet(
|
|
self, unet_name: str, weight_dtype: str, base_model: str = "Any"
|
|
) -> Tuple[Any, ...]:
|
|
"""Load a diffusion model by name, supporting extra folder paths
|
|
|
|
Args:
|
|
unet_name: The name of the diffusion model to load (relative path with extension)
|
|
weight_dtype: The dtype to use for model weights
|
|
base_model: Only used by the front-end to filter the random pool
|
|
|
|
Returns:
|
|
Tuple of (MODEL,)
|
|
"""
|
|
del base_model
|
|
import torch
|
|
|
|
# Get absolute path from cache using ComfyUI-style name
|
|
unet_path, metadata = get_checkpoint_info_absolute(unet_name)
|
|
|
|
if metadata is None:
|
|
raise FileNotFoundError(
|
|
f"Diffusion model '{unet_name}' not found in LoRA Manager cache. "
|
|
"Make sure the model is indexed and try again."
|
|
)
|
|
|
|
# Check if it's a GGUF model
|
|
if unet_path.endswith(".gguf"):
|
|
return self._load_gguf_unet(unet_path, unet_name, weight_dtype)
|
|
|
|
# Load regular diffusion model using ComfyUI's API
|
|
logger.info(f"Loading diffusion model from: {unet_path}")
|
|
|
|
# Build model options based on weight_dtype
|
|
model_options = {}
|
|
if weight_dtype == "fp8_e4m3fn":
|
|
model_options["dtype"] = torch.float8_e4m3fn
|
|
elif weight_dtype == "fp8_e4m3fn_fast":
|
|
model_options["dtype"] = torch.float8_e4m3fn
|
|
model_options["fp8_optimizations"] = True
|
|
elif weight_dtype == "fp8_e5m2":
|
|
model_options["dtype"] = torch.float8_e5m2
|
|
|
|
model = comfy.sd.load_diffusion_model(unet_path, model_options=model_options)
|
|
return (model,)
|
|
|
|
def _load_gguf_unet(
|
|
self, unet_path: str, unet_name: str, weight_dtype: str
|
|
) -> Tuple[Any, ...]:
|
|
"""Load a GGUF format diffusion model
|
|
|
|
Args:
|
|
unet_path: Absolute path to the GGUF file
|
|
unet_name: Name of the model for error messages
|
|
weight_dtype: The dtype to use for model weights
|
|
|
|
Returns:
|
|
Tuple of (MODEL,)
|
|
"""
|
|
import torch
|
|
from .gguf_import_helper import get_gguf_modules
|
|
|
|
# Get ComfyUI-GGUF modules using helper (handles various import scenarios)
|
|
try:
|
|
loader_module, ops_module, nodes_module = get_gguf_modules()
|
|
gguf_sd_loader = getattr(loader_module, "gguf_sd_loader")
|
|
GGMLOps = getattr(ops_module, "GGMLOps")
|
|
GGUFModelPatcher = getattr(nodes_module, "GGUFModelPatcher")
|
|
except RuntimeError as e:
|
|
raise RuntimeError(f"Cannot load GGUF model '{unet_name}'. {str(e)}")
|
|
|
|
logger.info(f"Loading GGUF diffusion model from: {unet_path}")
|
|
|
|
try:
|
|
# Load GGUF state dict
|
|
sd, extra = gguf_sd_loader(unet_path)
|
|
|
|
# Prepare kwargs for metadata if supported
|
|
kwargs = {}
|
|
import inspect
|
|
|
|
valid_params = inspect.signature(
|
|
comfy.sd.load_diffusion_model_state_dict
|
|
).parameters
|
|
if "metadata" in valid_params:
|
|
kwargs["metadata"] = extra.get("metadata", {})
|
|
|
|
# Setup custom operations with GGUF support
|
|
ops = GGMLOps()
|
|
|
|
# Handle weight_dtype for GGUF models
|
|
if weight_dtype in ("default", None):
|
|
ops.Linear.dequant_dtype = None
|
|
elif weight_dtype in ["target"]:
|
|
ops.Linear.dequant_dtype = weight_dtype
|
|
else:
|
|
ops.Linear.dequant_dtype = getattr(torch, weight_dtype, None)
|
|
|
|
# Load the model
|
|
model = comfy.sd.load_diffusion_model_state_dict(
|
|
sd, model_options={"custom_operations": ops}, **kwargs
|
|
)
|
|
|
|
if model is None:
|
|
raise RuntimeError(
|
|
f"Could not detect model type for GGUF diffusion model: {unet_path}"
|
|
)
|
|
|
|
# Wrap with GGUFModelPatcher
|
|
model = GGUFModelPatcher.clone(model)
|
|
|
|
# Register a reload factory so the MODEL carries its source path
|
|
# (cached_patcher_init) like core ComfyUI loaders do — required
|
|
# for model-name extraction downstream and for ModelPatcher
|
|
# deepclone/dynamic machinery.
|
|
model.cached_patcher_init = (_reload_gguf_unet, (unet_path, weight_dtype))
|
|
|
|
return (model,)
|
|
|
|
except Exception as e:
|
|
logger.error(f"Error loading GGUF diffusion model '{unet_name}': {e}")
|
|
raise RuntimeError(
|
|
f"Failed to load GGUF diffusion model '{unet_name}': {str(e)}"
|
|
)
|