mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-18 19:41:26 -03:00
feat(loaders): add control_after_generate random model selection to checkpoint/unet loaders
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).
This commit is contained in:
@@ -13,6 +13,10 @@ class CheckpointLoaderLM:
|
|||||||
|
|
||||||
Loads checkpoints from both standard ComfyUI folders and LoRA Manager's
|
Loads checkpoints from both standard ComfyUI folders and LoRA Manager's
|
||||||
extra folder paths, providing a unified interface for checkpoint loading.
|
extra folder paths, providing a unified interface for checkpoint loading.
|
||||||
|
The ckpt_name combo supports ComfyUI's control_after_generate, letting
|
||||||
|
users pick a random checkpoint on every run; the base_model input narrows
|
||||||
|
the random pool through a front-end extension that filters the combo
|
||||||
|
options.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
NAME = "Checkpoint Loader (LoraManager)"
|
NAME = "Checkpoint Loader (LoraManager)"
|
||||||
@@ -22,11 +26,29 @@ class CheckpointLoaderLM:
|
|||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
# Get list of checkpoint names from scanner (includes extra folder paths)
|
# Get list of checkpoint names from scanner (includes extra folder paths)
|
||||||
checkpoint_names = cls._get_checkpoint_names()
|
checkpoint_names = cls._get_checkpoint_names()
|
||||||
|
base_models = cls._get_available_base_models()
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"ckpt_name": (
|
"ckpt_name": (
|
||||||
checkpoint_names,
|
checkpoint_names,
|
||||||
{"tooltip": "The name of the checkpoint (model) to load."},
|
{
|
||||||
|
"tooltip": (
|
||||||
|
"The name of the checkpoint (model) to load. Use "
|
||||||
|
"control_after_generate to pick a random model on "
|
||||||
|
"every run."
|
||||||
|
),
|
||||||
|
"control_after_generate": True,
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"base_model": (
|
||||||
|
base_models,
|
||||||
|
{
|
||||||
|
"default": "Any",
|
||||||
|
"tooltip": (
|
||||||
|
"Restrict the random selection pool to this base "
|
||||||
|
"model. 'Any' uses the full pool."
|
||||||
|
),
|
||||||
|
},
|
||||||
),
|
),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -93,15 +115,68 @@ class CheckpointLoaderLM:
|
|||||||
logger.error(f"Error getting checkpoint names: {e}")
|
logger.error(f"Error getting checkpoint names: {e}")
|
||||||
return []
|
return []
|
||||||
|
|
||||||
def load_checkpoint(self, ckpt_name: str) -> Tuple[Any, Any, Any]:
|
@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"]
|
||||||
|
|
||||||
|
@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_checkpoint(
|
||||||
|
self, ckpt_name: str, base_model: str = "Any"
|
||||||
|
) -> Tuple[Any, Any, Any]:
|
||||||
"""Load a checkpoint by name, supporting extra folder paths
|
"""Load a checkpoint by name, supporting extra folder paths
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
ckpt_name: The name of the checkpoint to load (relative path with extension)
|
ckpt_name: The name of the checkpoint to load (relative path with extension)
|
||||||
|
base_model: Only used by the front-end to filter the random pool
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Tuple of (MODEL, CLIP, VAE)
|
Tuple of (MODEL, CLIP, VAE)
|
||||||
"""
|
"""
|
||||||
|
del base_model
|
||||||
# Get absolute path from cache using ComfyUI-style name
|
# Get absolute path from cache using ComfyUI-style name
|
||||||
ckpt_path, metadata = get_checkpoint_info_absolute(ckpt_name)
|
ckpt_path, metadata = get_checkpoint_info_absolute(ckpt_name)
|
||||||
|
|
||||||
|
|||||||
+77
-2
@@ -28,6 +28,10 @@ class UNETLoaderLM:
|
|||||||
Loads diffusion models/UNets from both standard ComfyUI folders and LoRA Manager's
|
Loads diffusion models/UNets from both standard ComfyUI folders and LoRA Manager's
|
||||||
extra folder paths, providing a unified interface for UNET loading.
|
extra folder paths, providing a unified interface for UNET loading.
|
||||||
Supports both regular diffusion models and GGUF format models.
|
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)"
|
NAME = "Unet Loader (LoraManager)"
|
||||||
@@ -37,16 +41,34 @@ class UNETLoaderLM:
|
|||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
# Get list of unet names from scanner (includes extra folder paths)
|
# Get list of unet names from scanner (includes extra folder paths)
|
||||||
unet_names = cls._get_unet_names()
|
unet_names = cls._get_unet_names()
|
||||||
|
base_models = cls._get_available_base_models()
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"unet_name": (
|
"unet_name": (
|
||||||
unet_names,
|
unet_names,
|
||||||
{"tooltip": "The name of the diffusion model to load."},
|
{
|
||||||
|
"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": (
|
"weight_dtype": (
|
||||||
["default", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e5m2"],
|
["default", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e5m2"],
|
||||||
{"tooltip": "The dtype to use for the model weights."},
|
{"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."
|
||||||
|
),
|
||||||
|
},
|
||||||
|
),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -108,16 +130,69 @@ class UNETLoaderLM:
|
|||||||
logger.error(f"Error getting unet names: {e}")
|
logger.error(f"Error getting unet names: {e}")
|
||||||
return []
|
return []
|
||||||
|
|
||||||
def load_unet(self, unet_name: str, weight_dtype: str) -> Tuple[Any, ...]:
|
@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
|
"""Load a diffusion model by name, supporting extra folder paths
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
unet_name: The name of the diffusion model to load (relative path with extension)
|
unet_name: The name of the diffusion model to load (relative path with extension)
|
||||||
weight_dtype: The dtype to use for model weights
|
weight_dtype: The dtype to use for model weights
|
||||||
|
base_model: Only used by the front-end to filter the random pool
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Tuple of (MODEL,)
|
Tuple of (MODEL,)
|
||||||
"""
|
"""
|
||||||
|
del base_model
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
# Get absolute path from cache using ComfyUI-style name
|
# Get absolute path from cache using ComfyUI-style name
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
import logging
|
import logging
|
||||||
|
import os
|
||||||
from typing import Any, Dict, List, Set
|
from typing import Any, Dict, List, Set
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
|
|
||||||
@@ -7,6 +8,7 @@ from .model_route_registrar import ModelRouteRegistrar
|
|||||||
from ..services.checkpoint_service import CheckpointService
|
from ..services.checkpoint_service import CheckpointService
|
||||||
from ..services.service_registry import ServiceRegistry
|
from ..services.service_registry import ServiceRegistry
|
||||||
from ..config import config
|
from ..config import config
|
||||||
|
from ..utils.utils import _format_model_name_for_comfyui
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -44,7 +46,45 @@ class CheckpointRoutes(BaseModelRoutes):
|
|||||||
# Checkpoint roots and Unet roots
|
# Checkpoint roots and Unet roots
|
||||||
registrar.add_prefixed_route('GET', '/api/lm/{prefix}/checkpoints_roots', prefix, self.get_checkpoints_roots)
|
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)
|
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
|
||||||
|
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.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
sub_type = request.query.get("sub_type", "checkpoint")
|
||||||
|
if sub_type not in ("checkpoint", "diffusion_model"):
|
||||||
|
return web.json_response({"error": "invalid sub_type"}, status=400)
|
||||||
|
scanner = await ServiceRegistry.get_checkpoint_scanner()
|
||||||
|
cache = await scanner.get_cached_data()
|
||||||
|
model_roots = scanner.get_model_roots()
|
||||||
|
items: List[Dict[str, str]] = []
|
||||||
|
for item in cache.raw_data:
|
||||||
|
if item.get("sub_type") != sub_type:
|
||||||
|
continue
|
||||||
|
file_path = item.get("file_path", "")
|
||||||
|
if not file_path or not os.path.exists(file_path):
|
||||||
|
continue
|
||||||
|
formatted_name = _format_model_name_for_comfyui(file_path, model_roots)
|
||||||
|
if formatted_name:
|
||||||
|
items.append(
|
||||||
|
{
|
||||||
|
"name": formatted_name,
|
||||||
|
"base_model": item.get("base_model", "") or "",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
items.sort(key=lambda x: x["name"])
|
||||||
|
return web.json_response({"items": items})
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error getting loader pool: {e}", exc_info=True)
|
||||||
|
return web.json_response({"error": str(e)}, status=500)
|
||||||
|
|
||||||
def _validate_civitai_model_type(self, model_type: str) -> bool:
|
def _validate_civitai_model_type(self, model_type: str) -> bool:
|
||||||
"""Validate CivitAI model type for Checkpoint"""
|
"""Validate CivitAI model type for Checkpoint"""
|
||||||
return model_type.lower() == 'checkpoint'
|
return model_type.lower() == 'checkpoint'
|
||||||
|
|||||||
@@ -83,3 +83,64 @@ def test_checkpoint_names_empty_when_scanner_fails(tmp_path, monkeypatch):
|
|||||||
|
|
||||||
monkeypatch.setattr(ServiceRegistry, "get_checkpoint_scanner", _boom)
|
monkeypatch.setattr(ServiceRegistry, "get_checkpoint_scanner", _boom)
|
||||||
assert CheckpointLoaderLM._get_checkpoint_names() == []
|
assert CheckpointLoaderLM._get_checkpoint_names() == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_checkpoint_available_base_models(tmp_path, monkeypatch):
|
||||||
|
from py.services.service_registry import ServiceRegistry
|
||||||
|
|
||||||
|
sd15 = tmp_path / "sd15.safetensors"
|
||||||
|
sd15.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(sd15), "base_model": "SD1.5"},
|
||||||
|
{"sub_type": "checkpoint", "file_path": str(flux), "base_model": "Flux.1 D"},
|
||||||
|
# Deleted files must drop out; wrong sub_type must be excluded.
|
||||||
|
{"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)
|
||||||
|
assert CheckpointLoaderLM._get_available_base_models() == [
|
||||||
|
"Any",
|
||||||
|
"Flux.1 D",
|
||||||
|
"SD1.5",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_unet_available_base_models(tmp_path, monkeypatch):
|
||||||
|
from py.services.service_registry import ServiceRegistry
|
||||||
|
|
||||||
|
flux = tmp_path / "flux.safetensors"
|
||||||
|
flux.write_bytes(b"x")
|
||||||
|
|
||||||
|
raw_data = [
|
||||||
|
{
|
||||||
|
"sub_type": "diffusion_model",
|
||||||
|
"file_path": str(flux),
|
||||||
|
"base_model": "Flux.1 D",
|
||||||
|
},
|
||||||
|
# Checkpoint entries must stay excluded by the sub_type filter.
|
||||||
|
{"sub_type": "checkpoint", "file_path": str(flux), "base_model": "SD1.5"},
|
||||||
|
]
|
||||||
|
|
||||||
|
async def _fake_scanner():
|
||||||
|
return _FakeScanner(raw_data, [str(tmp_path)])
|
||||||
|
|
||||||
|
monkeypatch.setattr(ServiceRegistry, "get_checkpoint_scanner", _fake_scanner)
|
||||||
|
assert UNETLoaderLM._get_available_base_models() == ["Any", "Flux.1 D"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_available_base_models_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 CheckpointLoaderLM._get_available_base_models() == ["Any"]
|
||||||
|
|||||||
@@ -0,0 +1,89 @@
|
|||||||
|
"""Tests for the loader-pool endpoint backing the Random Checkpoint/Unet
|
||||||
|
Loader nodes' front-end base_model filtering.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from py.routes.checkpoint_routes import CheckpointRoutes
|
||||||
|
from py.services.service_registry import ServiceRegistry
|
||||||
|
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
class DummyRequest:
|
||||||
|
def __init__(self, query=None):
|
||||||
|
self.query = query or {}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def routes(tmp_path, monkeypatch):
|
||||||
|
existing = tmp_path / "flux.safetensors"
|
||||||
|
existing.write_bytes(b"x")
|
||||||
|
missing = tmp_path / "missing.safetensors" # referenced but never created
|
||||||
|
|
||||||
|
raw_data = [
|
||||||
|
{"sub_type": "checkpoint", "file_path": str(existing), "base_model": "Flux.1 D"},
|
||||||
|
{"sub_type": "checkpoint", "file_path": str(missing), "base_model": "SDXL 1.0"},
|
||||||
|
{
|
||||||
|
"sub_type": "diffusion_model",
|
||||||
|
"file_path": str(existing),
|
||||||
|
"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 CheckpointRoutes()
|
||||||
|
|
||||||
|
|
||||||
|
async def test_loader_pool_checkpoint_subtype(routes):
|
||||||
|
response = await routes.get_loader_pool(DummyRequest(query={"sub_type": "checkpoint"}))
|
||||||
|
assert response.status == 200
|
||||||
|
payload = json.loads(response.text)
|
||||||
|
assert payload == {
|
||||||
|
"items": [{"name": "flux.safetensors", "base_model": "Flux.1 D"}]
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def test_loader_pool_diffusion_model_subtype(routes):
|
||||||
|
response = await routes.get_loader_pool(
|
||||||
|
DummyRequest(query={"sub_type": "diffusion_model"})
|
||||||
|
)
|
||||||
|
assert response.status == 200
|
||||||
|
payload = json.loads(response.text)
|
||||||
|
assert payload == {
|
||||||
|
"items": [{"name": "flux.safetensors", "base_model": "Flux.1 D"}]
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def test_loader_pool_default_subtype_is_checkpoint(routes):
|
||||||
|
response = await routes.get_loader_pool(DummyRequest())
|
||||||
|
assert response.status == 200
|
||||||
|
payload = json.loads(response.text)
|
||||||
|
assert payload == {
|
||||||
|
"items": [{"name": "flux.safetensors", "base_model": "Flux.1 D"}]
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def test_loader_pool_invalid_subtype(routes):
|
||||||
|
response = await routes.get_loader_pool(DummyRequest(query={"sub_type": "lora"}))
|
||||||
|
assert response.status == 400
|
||||||
@@ -0,0 +1,136 @@
|
|||||||
|
import { app } from "../../scripts/app.js";
|
||||||
|
import { api } from "../../scripts/api.js";
|
||||||
|
|
||||||
|
const NODE_CONFIGS = {
|
||||||
|
"Checkpoint Loader (LoraManager)": {
|
||||||
|
modelWidget: "ckpt_name",
|
||||||
|
subType: "checkpoint",
|
||||||
|
},
|
||||||
|
"Unet Loader (LoraManager)": {
|
||||||
|
modelWidget: "unet_name",
|
||||||
|
subType: "diffusion_model",
|
||||||
|
},
|
||||||
|
};
|
||||||
|
|
||||||
|
const poolCache = new Map();
|
||||||
|
|
||||||
|
async function fetchPool(subType) {
|
||||||
|
try {
|
||||||
|
const response = await api.fetchApi(
|
||||||
|
`/api/lm/checkpoints/loader-pool?sub_type=${encodeURIComponent(subType)}`
|
||||||
|
);
|
||||||
|
if (!response.ok) return [];
|
||||||
|
const data = await response.json();
|
||||||
|
return Array.isArray(data.items) ? data.items : [];
|
||||||
|
} catch (error) {
|
||||||
|
console.error("LoRA Manager: failed to fetch loader pool", error);
|
||||||
|
return [];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
async function refreshPoolCache() {
|
||||||
|
const subTypes = new Set(Object.values(NODE_CONFIGS).map((c) => c.subType));
|
||||||
|
await Promise.all(
|
||||||
|
[...subTypes].map(async (subType) => {
|
||||||
|
poolCache.set(subType, await fetchPool(subType));
|
||||||
|
})
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
function applyBaseModelFilter(node, config) {
|
||||||
|
const modelWidget = node.widgets?.find(
|
||||||
|
(widget) => widget.name === config.modelWidget
|
||||||
|
);
|
||||||
|
const baseModelWidget = node.widgets?.find(
|
||||||
|
(widget) => widget.name === "base_model"
|
||||||
|
);
|
||||||
|
if (!modelWidget || !baseModelWidget) return;
|
||||||
|
|
||||||
|
const wired = node.inputs?.some(
|
||||||
|
(input) =>
|
||||||
|
input.widget?.name === config.modelWidget && input.link != null
|
||||||
|
);
|
||||||
|
if (wired) return;
|
||||||
|
|
||||||
|
const pool = poolCache.get(config.subType) ?? [];
|
||||||
|
const filter = baseModelWidget.value;
|
||||||
|
const filtered =
|
||||||
|
filter === "Any"
|
||||||
|
? pool
|
||||||
|
: pool.filter((model) => model.base_model === filter);
|
||||||
|
const names = filtered.map((model) => model.name);
|
||||||
|
|
||||||
|
modelWidget.options.values = names;
|
||||||
|
if (!names.includes(modelWidget.value)) {
|
||||||
|
modelWidget.value = names[0];
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
function applyToAllNodes() {
|
||||||
|
app.graph?.nodes?.forEach((node) => {
|
||||||
|
const config = NODE_CONFIGS[node.comfyClass];
|
||||||
|
if (config) applyBaseModelFilter(node, config);
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
function ensureGraphConfigureHook(graph) {
|
||||||
|
if (!graph || graph.__loraManagerConfigureHooked) return;
|
||||||
|
graph.__loraManagerConfigureHooked = true;
|
||||||
|
|
||||||
|
const originalConfigure = graph.onConfigure;
|
||||||
|
graph.onConfigure = function (data) {
|
||||||
|
const result = originalConfigure?.call(this, data);
|
||||||
|
// Workflow reload restores widget values after onNodeCreated fires, so the
|
||||||
|
// per-node hook runs too early; re-apply the filter once the whole graph
|
||||||
|
// has been configured.
|
||||||
|
setTimeout(() => applyToAllNodes(), 0);
|
||||||
|
return result;
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
app.registerExtension({
|
||||||
|
name: "LoraManager.RandomLoaderControl",
|
||||||
|
|
||||||
|
async setup() {
|
||||||
|
await refreshPoolCache();
|
||||||
|
},
|
||||||
|
|
||||||
|
beforeRegisterNodeDef(nodeType, nodeData) {
|
||||||
|
const config = NODE_CONFIGS[nodeType.comfyClass];
|
||||||
|
if (!config) return;
|
||||||
|
|
||||||
|
const onNodeCreated = nodeType.prototype.onNodeCreated;
|
||||||
|
nodeType.prototype.onNodeCreated = function () {
|
||||||
|
const result = onNodeCreated?.apply(this, arguments);
|
||||||
|
|
||||||
|
const baseModelWidget = this.widgets?.find(
|
||||||
|
(widget) => widget.name === "base_model"
|
||||||
|
);
|
||||||
|
if (baseModelWidget) {
|
||||||
|
const originalCallback = baseModelWidget.callback;
|
||||||
|
baseModelWidget.callback = (value, canvas, node, pos, event) => {
|
||||||
|
applyBaseModelFilter(node ?? this, config);
|
||||||
|
return originalCallback?.call(this, value, canvas, node, pos, event);
|
||||||
|
};
|
||||||
|
}
|
||||||
|
|
||||||
|
applyBaseModelFilter(this, config);
|
||||||
|
return result;
|
||||||
|
};
|
||||||
|
|
||||||
|
// onNodeCreated fires inside LGraph.createNode, before the node is added to
|
||||||
|
// a graph (this.graph is null there), so the graph-level configure hook
|
||||||
|
// must be installed from onAdded, where the graph reference is available.
|
||||||
|
const onAdded = nodeType.prototype.onAdded;
|
||||||
|
nodeType.prototype.onAdded = function () {
|
||||||
|
const result = onAdded?.apply(this, arguments);
|
||||||
|
ensureGraphConfigureHook(this.graph);
|
||||||
|
return result;
|
||||||
|
};
|
||||||
|
},
|
||||||
|
|
||||||
|
async refreshComboInNodes() {
|
||||||
|
await refreshPoolCache();
|
||||||
|
applyToAllNodes();
|
||||||
|
},
|
||||||
|
});
|
||||||
Reference in New Issue
Block a user