Handle all the model lists for the selectors in one function, and cache the results for a minute.

This commit is contained in:
Shanoah Alkire
2025-09-29 00:15:40 -07:00
parent f8c3b95a43
commit f6312e9d6c
3 changed files with 63 additions and 48 deletions
+13 -47
View File
@@ -8,6 +8,7 @@ import folder_paths
from ..utils import model_info as mi
from comfy_execution.graph_utils import GraphBuilder
from ..utils import add_lora_to_stack
from ..utils import get_model_list
from ..utils.helpers_graph import (
add_lora_stack_node
)
@@ -16,7 +17,7 @@ from ..utils.helpers_graph import (
class Sage_CheckpointSelector(ComfyNodeABC):
@classmethod
def INPUT_TYPES(cls) -> InputTypeDict:
model_list = folder_paths.get_filename_list("checkpoints")
model_list = get_model_list("checkpoints")
return {
"required": {
"ckpt_name": (model_list, {"tooltip": "The name of the checkpoint (model) to load."})
@@ -39,14 +40,7 @@ class Sage_CheckpointSelector(ComfyNodeABC):
class Sage_UNETSelector(ComfyNodeABC):
@classmethod
def INPUT_TYPES(cls) -> InputTypeDict:
unet_names = []
try:
unet_names = [x for x in folder_paths.get_filename_list("unet_gguf")]
except Exception as e:
unet_names = []
unet_names += folder_paths.get_filename_list("diffusion_models")
unet_names = list(set(unet_names))
unet_names.sort() # Remove duplicates
unet_names = get_model_list("unet")
return {
"required": {
"unet_name": (unet_names, {"tooltip": "The name of the UNET model to load."}),
@@ -57,7 +51,7 @@ class Sage_UNETSelector(ComfyNodeABC):
RETURN_TYPES = ("UNET_INFO",)
RETURN_NAMES = ("unet_info",)
OUTPUT_TOOLTIPS = ("The model path and hash, all in one output.")
OUTPUT_TOOLTIPS = ("The model path and hash, all in one output.",)
FUNCTION = "get_unet_info"
CATEGORY = "Sage Utils/selectors"
@@ -69,7 +63,7 @@ class Sage_UNETSelector(ComfyNodeABC):
class Sage_VAESelector(ComfyNodeABC):
@classmethod
def INPUT_TYPES(cls) -> InputTypeDict:
vae_list = folder_paths.get_filename_list("vae")
vae_list = get_model_list("vae")
return {
"required": {
"vae_name": (vae_list, {"tooltip": "The name of the VAE model to load."})
@@ -91,14 +85,7 @@ class Sage_VAESelector(ComfyNodeABC):
class Sage_CLIPSelector(ComfyNodeABC):
@classmethod
def INPUT_TYPES(cls) -> InputTypeDict:
model_list = []
try:
model_list = [x for x in folder_paths.get_filename_list("clip_gguf")]
except Exception as e:
model_list = []
model_list += folder_paths.get_filename_list("text_encoders")
model_list = list(set(model_list))
model_list.sort() # Remove duplicates
model_list = get_model_list("clip")
return {
"required": {
"clip_name": (model_list, {"tooltip": "The name of the CLIP model to load."}),
@@ -121,14 +108,7 @@ class Sage_CLIPSelector(ComfyNodeABC):
class Sage_DualCLIPSelector(ComfyNodeABC):
@classmethod
def INPUT_TYPES(cls) -> InputTypeDict:
model_list = []
try:
model_list = [x for x in folder_paths.get_filename_list("clip_gguf")]
except Exception as e:
model_list = []
model_list += folder_paths.get_filename_list("text_encoders")
model_list = list(set(model_list))
model_list.sort() # Remove duplicates
model_list = get_model_list("clip")
return {
"required": {
"clip_name_1": (model_list, {"tooltip": "The name of the first CLIP model to load."}),
@@ -152,14 +132,7 @@ class Sage_DualCLIPSelector(ComfyNodeABC):
class Sage_TripleCLIPSelector(ComfyNodeABC):
@classmethod
def INPUT_TYPES(cls) -> InputTypeDict:
model_list = []
try:
model_list = [x for x in folder_paths.get_filename_list("clip_gguf")]
except Exception as e:
model_list = []
model_list += folder_paths.get_filename_list("text_encoders")
model_list = list(set(model_list))
model_list.sort() # Remove duplicates
model_list = get_model_list("clip")
return {
"required": {
"clip_name_1": (model_list, {"tooltip": "The name of the first CLIP model to load."}),
@@ -182,14 +155,7 @@ class Sage_TripleCLIPSelector(ComfyNodeABC):
class Sage_QuadCLIPSelector(ComfyNodeABC):
@classmethod
def INPUT_TYPES(cls) -> InputTypeDict:
model_list = []
try:
model_list = [x for x in folder_paths.get_filename_list("clip_gguf")]
except Exception as e:
model_list = []
model_list += folder_paths.get_filename_list("text_encoders")
model_list = list(set(model_list))
model_list.sort() # Remove duplicates
model_list = get_model_list("clip")
return {
"required": {
"clip_name_1": (model_list, {"tooltip": "The name of the first CLIP model to load."}),
@@ -330,7 +296,7 @@ class Sage_LoraStack(ComfyNodeABC):
@classmethod
def INPUT_TYPES(cls) -> InputTypeDict:
lora_list = folder_paths.get_filename_list("loras")
lora_list = get_model_list("loras")
return {
"required": {
"enabled": (IO.BOOLEAN, {"default": False, "tooltip": "Whether to enable this LoRA."}),
@@ -366,7 +332,7 @@ class Sage_QuickLoraStack(Sage_LoraStack):
@classmethod
def INPUT_TYPES(cls) -> InputTypeDict:
lora_list = folder_paths.get_filename_list("loras")
lora_list = get_model_list("loras")
return {
"required": {
"enabled": (IO.BOOLEAN, {"default": True}),
@@ -397,7 +363,7 @@ class Sage_TripleLoraStack(ComfyNodeABC):
@classmethod
def INPUT_TYPES(cls) -> InputTypeDict:
lora_list = folder_paths.get_filename_list("loras")
lora_list = get_model_list("loras")
required_list = {}
for i in range(1, cls.NUM_OF_ENTRIES + 1):
required_list[f"enabled_{i}"] = (IO.BOOLEAN, {"default": True})
@@ -448,7 +414,7 @@ class Sage_TripleQuickLoraStack(ComfyNodeABC):
@classmethod
def INPUT_TYPES(cls) -> InputTypeDict:
lora_list = folder_paths.get_filename_list("loras")
lora_list = get_model_list("loras")
required_list = {}
for i in range(1, cls.NUM_OF_ENTRIES + 1):
required_list[f"enabled_{i}"] = (IO.BOOLEAN, {"default": True})
+2 -1
View File
@@ -26,6 +26,7 @@ from .helpers import (
lora_to_prompt,
get_lora_hash,
model_scan,
get_model_list,
get_recently_used_models,
clean_keywords,
clean_text,
@@ -125,7 +126,7 @@ __all__ = [
'get_files_in_dir', 'last_used', 'days_since_last_used', 'get_file_modification_date',
'update_cache_from_civitai_json', 'update_cache_without_civitai_json',
'add_file_to_cache', 'recheck_hash', 'pull_metadata', 'update_model_timestamp', 'pull_and_update_model_timestamp',
'lora_to_string', 'lora_to_prompt', 'get_lora_hash', 'model_scan',
'lora_to_string', 'lora_to_prompt', 'get_lora_hash', 'model_scan', 'get_model_list',
'get_recently_used_models', 'clean_keywords', 'clean_text', 'condition_text',
'get_save_file_path', 'unwrap_tuple',
+48
View File
@@ -433,6 +433,54 @@ def model_scan(the_path, force = False):
pbar = comfy.utils.ProgressBar(len(model_list))
pull_metadata(model_list, force_all=force, pbar=pbar)
# Module-level cache for get_model_list
_model_list_cache = {}
def get_model_list(model_type: str) -> list[str]:
"""Get a list of model names based on the model type, with in-memory cache (1 min)."""
import time
global _model_list_cache
now = time.time()
cache_entry = _model_list_cache.get(model_type)
if cache_entry:
cached_time, cached_list = cache_entry
if now - cached_time < 60:
return cached_list
# Not cached or cache expired, fetch fresh
if model_type == "checkpoints":
result = folder_paths.get_filename_list("checkpoints")
elif model_type == "unet":
unet_names = []
try:
unet_names = [x for x in folder_paths.get_filename_list("unet_gguf")]
except Exception as e:
unet_names = []
unet_names += folder_paths.get_filename_list("diffusion_models")
unet_names = list(set(unet_names))
unet_names.sort() # Remove duplicates
result = unet_names
elif model_type == "vae":
result = folder_paths.get_filename_list("vae")
elif model_type == "clip":
model_list = []
try:
model_list = [x for x in folder_paths.get_filename_list("clip_gguf")]
except Exception as e:
model_list = []
model_list += folder_paths.get_filename_list("text_encoders")
model_list = list(set(model_list))
model_list.sort() # Remove duplicates
result = model_list
elif model_type == "loras":
result = folder_paths.get_filename_list("loras")
else:
result = []
_model_list_cache[model_type] = (now, result)
return result
def get_recently_used_models(model_type):
model_list = list()
full_model_list = folder_paths.get_filename_list(model_type)