Handle all the model lists for the selectors in one function, and cache the results for a minute.
This commit is contained in:
+13
-47
@@ -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
@@ -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',
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user