From 3fd29f69439929f8f24631b12dd7e6dacc5d6e4a Mon Sep 17 00:00:00 2001 From: Will Miao Date: Sun, 30 Aug 2026 11:38:01 +0800 Subject: [PATCH] remove(nodes): delete Random Checkpoint Loader and Random Unet Loader nodes - Remove py/nodes/random_checkpoint_loader.py and random_unet_loader.py - Remove their dedicated test file - Clean up imports and NODE_CLASS_MAPPINGS in __init__.py - Update loader-pool comments/docstrings to reference the remaining Checkpoint/Unet Loader nodes' control_after_generate feature --- __init__.py | 10 - py/nodes/random_checkpoint_loader.py | 214 ------------ py/nodes/random_unet_loader.py | 326 ------------------ py/routes/checkpoint_routes.py | 9 +- .../test_random_checkpoint_unet_loaders.py | 179 ---------- tests/routes/test_random_loader_pool.py | 4 +- 6 files changed, 7 insertions(+), 735 deletions(-) delete mode 100644 py/nodes/random_checkpoint_loader.py delete mode 100644 py/nodes/random_unet_loader.py delete mode 100644 tests/nodes/test_random_checkpoint_unet_loaders.py diff --git a/__init__.py b/__init__.py index 4f9b0989..a5b691bc 100644 --- a/__init__.py +++ b/__init__.py @@ -3,8 +3,6 @@ try: # pragma: no cover - import fallback for pytest collection from .py.nodes.lora_loader import LoraLoaderLM, LoraTextLoaderLM from .py.nodes.checkpoint_loader import CheckpointLoaderLM from .py.nodes.unet_loader import UNETLoaderLM - from .py.nodes.random_checkpoint_loader import RandomCheckpointLoaderLM - from .py.nodes.random_unet_loader import RandomUNETLoaderLM from .py.nodes.trigger_word_toggle import TriggerWordToggleLM from .py.nodes.prompt import PromptLM from .py.nodes.text import TextLM @@ -42,12 +40,6 @@ except ( "py.nodes.checkpoint_loader" ).CheckpointLoaderLM UNETLoaderLM = importlib.import_module("py.nodes.unet_loader").UNETLoaderLM - RandomCheckpointLoaderLM = importlib.import_module( - "py.nodes.random_checkpoint_loader" - ).RandomCheckpointLoaderLM - RandomUNETLoaderLM = importlib.import_module( - "py.nodes.random_unet_loader" - ).RandomUNETLoaderLM TriggerWordToggleLM = importlib.import_module( "py.nodes.trigger_word_toggle" ).TriggerWordToggleLM @@ -87,8 +79,6 @@ NODE_CLASS_MAPPINGS = { LoraTextLoaderLM.NAME: LoraTextLoaderLM, CheckpointLoaderLM.NAME: CheckpointLoaderLM, UNETLoaderLM.NAME: UNETLoaderLM, - RandomCheckpointLoaderLM.NAME: RandomCheckpointLoaderLM, - RandomUNETLoaderLM.NAME: RandomUNETLoaderLM, TriggerWordToggleLM.NAME: TriggerWordToggleLM, LoraStackerLM.NAME: LoraStackerLM, LoraStackCombinerLM.NAME: LoraStackCombinerLM, diff --git a/py/nodes/random_checkpoint_loader.py b/py/nodes/random_checkpoint_loader.py deleted file mode 100644 index b9c4ff24..00000000 --- a/py/nodes/random_checkpoint_loader.py +++ /dev/null @@ -1,214 +0,0 @@ -import logging -import os -import random -from typing import Any, List, Optional, Tuple -import comfy.sd # pyright: ignore[reportMissingImports] -import folder_paths # pyright: ignore[reportMissingImports] -from ..utils.utils import get_checkpoint_info_absolute, _format_model_name_for_comfyui - -logger = logging.getLogger(__name__) - - -class RandomCheckpointLoaderLM: - """Checkpoint Loader that can randomly pick a checkpoint from the pool - - Loads checkpoints from both standard ComfyUI folders and LoRA Manager's - extra folder paths. When select_at_random is enabled, ignores ckpt_name - and picks a random checkpoint (optionally filtered by base_model) on - every run. - """ - - NAME = "Random Checkpoint Loader (LoraManager)" - CATEGORY = "Lora Manager/loaders" - - @classmethod - def INPUT_TYPES(cls): - # Get list of checkpoint names from scanner (includes extra folder paths) - checkpoint_names = cls._get_checkpoint_names() - base_models = cls._get_available_base_models() - return { - "required": { - "ckpt_name": ( - checkpoint_names, - {"tooltip": "The name of the checkpoint (model) to load."}, - ), - "select_at_random": ( - "BOOLEAN", - { - "default": False, - "tooltip": ( - "Ignore ckpt_name and pick a random checkpoint 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", "CLIP", "VAE", "STRING") - RETURN_NAMES = ("MODEL", "CLIP", "VAE", "model_name") - OUTPUT_TOOLTIPS = ( - "The model used for denoising latents.", - "The CLIP model used for encoding text prompts.", - "The VAE model used for encoding and decoding images to and from latent space.", - "The name of the checkpoint that was loaded (useful when select_at_random is enabled).", - ) - FUNCTION = "load_checkpoint" - - @classmethod - def IS_CHANGED(cls, ckpt_name, 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 ckpt_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_checkpoint_names(cls, base_model: Optional[str] = None) -> List[str]: - """Get list of checkpoint names from scanner cache in ComfyUI format (relative path with extension) - - Args: - base_model: If given (and not "Any"), only include checkpoints 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 checkpoint type (not diffusion_model) and format names - names = [] - for item in cache.raw_data: - if item.get("sub_type") != "checkpoint": - 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 checkpoints 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 checkpoint names: {e}") - return [] - - @classmethod - def _get_available_base_models(cls) -> List[str]: - """Get distinct base_model values present among indexed checkpoints, 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") != "checkpoint": - 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_checkpoint( - self, - ckpt_name: str, - select_at_random: bool = False, - base_model: str = "Any", - ) -> Tuple[Any, Any, Any, str]: - """Load a checkpoint by name, supporting extra folder paths - - Args: - ckpt_name: The name of the checkpoint to load (relative path with extension) - select_at_random: If True, ignore ckpt_name and pick randomly from the pool - base_model: Restricts random selection to this base model ("Any" = no filter) - - Returns: - Tuple of (MODEL, CLIP, VAE, model_name) - """ - if select_at_random: - pool = self._get_checkpoint_names(base_model) - if not pool: - raise FileNotFoundError( - f"No checkpoints found for base model '{base_model}'. " - "Pick a different base model or disable 'select_at_random'." - ) - ckpt_name = random.choice(pool) - logger.info( - f"[RandomCheckpointLoaderLM] Randomly selected checkpoint: {ckpt_name}" - ) - - # Get absolute path from cache using ComfyUI-style name - ckpt_path, metadata = get_checkpoint_info_absolute(ckpt_name) - - if metadata is None: - raise FileNotFoundError( - f"Checkpoint '{ckpt_name}' not found in LoRA Manager cache. " - "Make sure the checkpoint is indexed and try again." - ) - - # Load regular checkpoint using ComfyUI's API - logger.info(f"Loading checkpoint from: {ckpt_path}") - out = comfy.sd.load_checkpoint_guess_config( - ckpt_path, - output_vae=True, - output_clip=True, - embedding_directory=folder_paths.get_folder_paths("embeddings"), - ) - return out[:3] + (ckpt_name,) diff --git a/py/nodes/random_unet_loader.py b/py/nodes/random_unet_loader.py deleted file mode 100644 index f18489a2..00000000 --- a/py/nodes/random_unet_loader.py +++ /dev/null @@ -1,326 +0,0 @@ -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)}" - ) diff --git a/py/routes/checkpoint_routes.py b/py/routes/checkpoint_routes.py index 4f588a83..0187f856 100644 --- a/py/routes/checkpoint_routes.py +++ b/py/routes/checkpoint_routes.py @@ -47,15 +47,16 @@ class CheckpointRoutes(BaseModelRoutes): registrar.add_prefixed_route('GET', '/api/lm/{prefix}/checkpoints_roots', prefix, self.get_checkpoints_roots) registrar.add_prefixed_route('GET', '/api/lm/{prefix}/unet_roots', prefix, self.get_unet_roots) - # Name/base_model pool for the Random Checkpoint/Unet Loader nodes + # Name/base_model pool for the Checkpoint/Unet Loader nodes' base_model filtering registrar.add_prefixed_route('GET', '/api/lm/{prefix}/loader-pool', prefix, self.get_loader_pool) async def get_loader_pool(self, request: web.Request) -> web.Response: """Return ComfyUI-formatted model names with their base_model. - Backing data for the Random Checkpoint/Unet Loader nodes: the front-end - filters the ckpt_name/unet_name combo options by base_model using this - pool, so control_after_generate randomizes within the narrowed set. + Backing data for the Checkpoint/Unet Loader nodes' + control_after_generate feature: the front-end filters the + ckpt_name/unet_name combo options by base_model using this pool, so + randomize mode picks within the narrowed set. """ try: sub_type = request.query.get("sub_type", "checkpoint") diff --git a/tests/nodes/test_random_checkpoint_unet_loaders.py b/tests/nodes/test_random_checkpoint_unet_loaders.py deleted file mode 100644 index 362ad434..00000000 --- a/tests/nodes/test_random_checkpoint_unet_loaders.py +++ /dev/null @@ -1,179 +0,0 @@ -"""Tests for the Random Checkpoint/Unet Loader nodes' base-model filtering and -random-selection behavior. -""" - -import pytest - -from py.nodes.random_checkpoint_loader import RandomCheckpointLoaderLM -from py.nodes.random_unet_loader import RandomUNETLoaderLM - - -class _FakeCache: - def __init__(self, raw_data): - self.raw_data = raw_data - - -class _FakeScanner: - def __init__(self, raw_data, model_roots): - self._raw_data = raw_data - self._model_roots = model_roots - - async def get_cached_data(self, force_refresh=False): - return _FakeCache(self._raw_data) - - def get_model_roots(self): - return self._model_roots - - -@pytest.fixture -def base_model_library(tmp_path, monkeypatch): - from py.services.service_registry import ServiceRegistry - - illustrious = tmp_path / "illustrious.safetensors" - illustrious.write_bytes(b"x") - flux = tmp_path / "flux.safetensors" - flux.write_bytes(b"x") - missing = tmp_path / "missing.safetensors" # referenced but never created - - raw_data = [ - { - "sub_type": "checkpoint", - "file_path": str(illustrious), - "base_model": "Illustrious", - }, - {"sub_type": "checkpoint", "file_path": str(flux), "base_model": "Flux.1 D"}, - { - "sub_type": "checkpoint", - "file_path": str(missing), - "base_model": "SDXL 1.0", - }, - { - "sub_type": "diffusion_model", - "file_path": str(flux), - "base_model": "Flux.1 D", - }, - ] - - async def _fake_scanner(): - return _FakeScanner(raw_data, [str(tmp_path)]) - - monkeypatch.setattr(ServiceRegistry, "get_checkpoint_scanner", _fake_scanner) - return tmp_path - - -def test_checkpoint_names_drop_deleted_files(tmp_path, monkeypatch): - from py.services.service_registry import ServiceRegistry - - existing = tmp_path / "keep.safetensors" - existing.write_bytes(b"x") - deleted = tmp_path / "deleted.safetensors" # referenced but never created - - raw_data = [ - {"sub_type": "checkpoint", "file_path": str(existing)}, - {"sub_type": "checkpoint", "file_path": str(deleted)}, - # Wrong type must stay excluded by the sub_type filter. - {"sub_type": "diffusion_model", "file_path": str(existing)}, - ] - - async def _fake_scanner(): - return _FakeScanner(raw_data, [str(tmp_path)]) - - monkeypatch.setattr(ServiceRegistry, "get_checkpoint_scanner", _fake_scanner) - assert RandomCheckpointLoaderLM._get_checkpoint_names() == ["keep.safetensors"] - - -def test_unet_names_drop_deleted_files(tmp_path, monkeypatch): - from py.services.service_registry import ServiceRegistry - - existing = tmp_path / "keep.safetensors" - existing.write_bytes(b"x") - deleted = tmp_path / "deleted.safetensors" - - raw_data = [ - {"sub_type": "diffusion_model", "file_path": str(existing)}, - {"sub_type": "diffusion_model", "file_path": str(deleted)}, - {"sub_type": "checkpoint", "file_path": str(existing)}, - ] - - async def _fake_scanner(): - return _FakeScanner(raw_data, [str(tmp_path)]) - - monkeypatch.setattr(ServiceRegistry, "get_checkpoint_scanner", _fake_scanner) - assert RandomUNETLoaderLM._get_unet_names() == ["keep.safetensors"] - - -def test_checkpoint_names_empty_when_scanner_fails(tmp_path, monkeypatch): - from py.services.service_registry import ServiceRegistry - - def _boom(): - raise RuntimeError("scanner not available") - - monkeypatch.setattr(ServiceRegistry, "get_checkpoint_scanner", _boom) - assert RandomCheckpointLoaderLM._get_checkpoint_names() == [] - - -def test_checkpoint_available_base_models(base_model_library): - # "SDXL 1.0" is excluded because its file no longer exists on disk. - assert RandomCheckpointLoaderLM._get_available_base_models() == [ - "Any", - "Flux.1 D", - "Illustrious", - ] - - -def test_checkpoint_names_filtered_by_base_model(base_model_library): - assert RandomCheckpointLoaderLM._get_checkpoint_names("Illustrious") == [ - "illustrious.safetensors" - ] - assert RandomCheckpointLoaderLM._get_checkpoint_names("Any") == [ - "flux.safetensors", - "illustrious.safetensors", - ] - - -def test_unet_available_base_models(base_model_library): - assert RandomUNETLoaderLM._get_available_base_models() == ["Any", "Flux.1 D"] - - -def test_load_checkpoint_random_selection_uses_pool(base_model_library, monkeypatch): - from py.nodes import random_checkpoint_loader as random_checkpoint_loader_module - - monkeypatch.setattr( - random_checkpoint_loader_module, - "get_checkpoint_info_absolute", - lambda name: (str(base_model_library / name), {"file_path": name}), - ) - monkeypatch.setattr( - random_checkpoint_loader_module.comfy.sd, - "load_checkpoint_guess_config", - lambda *a, **k: ("MODEL", "CLIP", "VAE", None), - raising=False, - ) - - node = RandomCheckpointLoaderLM() - result = node.load_checkpoint( - "ignored.safetensors", select_at_random=True, base_model="Illustrious" - ) - # Only one checkpoint matches "Illustrious", so the random pick is deterministic here. - assert result[3] == "illustrious.safetensors" - - -def test_load_checkpoint_random_selection_raises_when_pool_empty(base_model_library): - node = RandomCheckpointLoaderLM() - with pytest.raises(FileNotFoundError, match="No checkpoints found"): - node.load_checkpoint( - "ignored.safetensors", select_at_random=True, base_model="SDXL 1.0" - ) - - -def test_checkpoint_is_changed_forces_rerun_when_random(): - assert RandomCheckpointLoaderLM.IS_CHANGED( - "a.safetensors", select_at_random=True, base_model="Any" - ) != RandomCheckpointLoaderLM.IS_CHANGED( - "a.safetensors", select_at_random=True, base_model="Any" - ) - assert RandomCheckpointLoaderLM.IS_CHANGED( - "a.safetensors", select_at_random=False, base_model="Any" - ) == RandomCheckpointLoaderLM.IS_CHANGED( - "a.safetensors", select_at_random=False, base_model="Any" - ) diff --git a/tests/routes/test_random_loader_pool.py b/tests/routes/test_random_loader_pool.py index cbc6f662..ecd424b6 100644 --- a/tests/routes/test_random_loader_pool.py +++ b/tests/routes/test_random_loader_pool.py @@ -1,5 +1,5 @@ -"""Tests for the loader-pool endpoint backing the Random Checkpoint/Unet -Loader nodes' front-end base_model filtering. +"""Tests for the loader-pool endpoint backing the Checkpoint/Unet Loader +nodes' front-end base_model filtering. """ import json