import logging import os import random from typing import Any, List, Optional, 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 RandomUNETLoaderLM.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 = RandomUNETLoaderLM() model, _unet_name = loader._load_gguf_unet(unet_path, unet_path, weight_dtype) return model class RandomUNETLoaderLM: """UNET Loader that can randomly pick a diffusion model from the pool Loads diffusion models/UNets from both standard ComfyUI folders and LoRA Manager's extra folder paths. Supports both regular diffusion models and GGUF format models. When select_at_random is enabled, ignores unet_name and picks a random diffusion model (optionally filtered by base_model) on every run. """ NAME = "Random 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."}, ), "weight_dtype": ( ["default", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e5m2"], {"tooltip": "The dtype to use for the model weights."}, ), "select_at_random": ( "BOOLEAN", { "default": False, "tooltip": ( "Ignore unet_name and pick a random diffusion model from " "the pool (optionally filtered by base_model) on every run." ), }, ), "base_model": ( base_models, { "default": "Any", "tooltip": "Restrict random selection to this base model. 'Any' uses the full pool.", }, ), } } RETURN_TYPES = ("MODEL", "STRING") RETURN_NAMES = ("MODEL", "model_name") OUTPUT_TOOLTIPS = ( "The model used for denoising latents.", "The name of the diffusion model that was loaded (useful when select_at_random is enabled).", ) FUNCTION = "load_unet" @classmethod def IS_CHANGED( cls, unet_name, weight_dtype, select_at_random=False, base_model="Any" ): # Force re-execution on every run while randomizing, since the widget # values themselves don't change between queue runs. if select_at_random: return float("nan") return unet_name @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()) @classmethod def _get_unet_names(cls, base_model: Optional[str] = None) -> List[str]: """Get list of diffusion model names from scanner cache in ComfyUI format (relative path with extension) Args: base_model: If given (and not "Any"), only include models matching this base model. """ try: from ..services.service_registry import ServiceRegistry 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": continue if ( base_model and base_model != "Any" and item.get("base_model") != base_model ): continue 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) return cls._run_async(_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"] def load_unet( self, unet_name: str, weight_dtype: str, select_at_random: bool = False, 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 select_at_random: If True, ignore unet_name and pick randomly from the pool base_model: Restricts random selection to this base model ("Any" = no filter) Returns: Tuple of (MODEL, model_name) """ import torch if select_at_random: pool = self._get_unet_names(base_model) if not pool: raise FileNotFoundError( f"No diffusion models found for base model '{base_model}'. " "Pick a different base model or disable 'select_at_random'." ) unet_name = random.choice(pool) logger.info( f"[RandomUNETLoaderLM] Randomly selected diffusion model: {unet_name}" ) # 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, unet_name) 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, model_name) """ 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, unet_name) 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)}" )