From 98cba08d578de940e3b2892f09e030ff640f8ccc Mon Sep 17 00:00:00 2001 From: Haoming Date: Wed, 6 Aug 2025 17:47:21 +0800 Subject: [PATCH] ckpt --- modules/extras.py | 4 +--- modules/sd_models.py | 18 +++++++++--------- modules/sd_vae.py | 7 ------- modules/ui_extra_networks_checkpoints.py | 2 +- 4 files changed, 11 insertions(+), 20 deletions(-) diff --git a/modules/extras.py b/modules/extras.py index b13dab43..365457cb 100644 --- a/modules/extras.py +++ b/modules/extras.py @@ -215,14 +215,12 @@ def run_modelmerger(id_task, primary_model_name, secondary_model_name, tertiary_ if re.search(regex, key): theta_0.pop(key, None) - ckpt_dir = shared.cmd_opts.ckpt_dir or sd_models.model_path - filename = filename_generator() if custom_name == '' else custom_name filename += ".inpainting" if result_is_inpainting_model else "" filename += ".instruct-pix2pix" if result_is_instruct_pix2pix_model else "" filename += "." + checkpoint_format - output_modelname = os.path.join(ckpt_dir, filename) + output_modelname = os.path.join(sd_models.model_path, filename) shared.state.nextjob() shared.state.textinfo = "Saving" diff --git a/modules/sd_models.py b/modules/sd_models.py index 693486af..267b8f33 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -63,16 +63,14 @@ class CheckpointInfo: def __init__(self, filename): self.filename = filename abspath = os.path.abspath(filename) - abs_ckpt_dir = os.path.abspath(shared.cmd_opts.ckpt_dir) if shared.cmd_opts.ckpt_dir is not None else None + abs_ckpt_dirs = (*cmd_opts.ckpt_dirs, model_path) self.is_safetensors = os.path.splitext(filename)[1].lower() == ".safetensors" - if abs_ckpt_dir and abspath.startswith(abs_ckpt_dir): - name = abspath.replace(abs_ckpt_dir, "") - elif abspath.startswith(model_path): - name = abspath.replace(model_path, "") - else: - name = os.path.basename(filename) + for _dir in abs_ckpt_dirs: + if abspath.startswith(str(_dir)): + name = abspath.replace(str(_dir), "") + break if name.startswith("\\") or name.startswith("/"): name = name[1:] @@ -156,10 +154,12 @@ def checkpoint_tiles(use_short=False): def list_models(): checkpoints_list.clear() checkpoint_aliases.clear() + model_list = [] - model_list = modelloader.load_models(model_path=model_path, model_url=None, command_path=shared.cmd_opts.ckpt_dir, ext_filter=[".ckpt", ".safetensors", ".gguf"], download_name=None, ext_blacklist=[".vae.ckpt", ".vae.safetensors"]) + for _dir in (*cmd_opts.ckpt_dirs, model_path): + model_list.extend(modelloader.load_models(model_path=_dir, ext_filter=[".ckpt", ".safetensors", ".gguf"], ext_blacklist=[".vae.ckpt", ".vae.safetensors"])) - for filename in model_list: + for filename in set(model_list): checkpoint_info = CheckpointInfo(filename) checkpoint_info.register() diff --git a/modules/sd_vae.py b/modules/sd_vae.py index 2943fbf4..443149fc 100644 --- a/modules/sd_vae.py +++ b/modules/sd_vae.py @@ -82,13 +82,6 @@ def refresh_vae_list(): os.path.join(vae_path, '**/*.safetensors'), ] - if shared.cmd_opts.ckpt_dir is not None and os.path.isdir(shared.cmd_opts.ckpt_dir): - paths += [ - os.path.join(shared.cmd_opts.ckpt_dir, '**/*.vae.ckpt'), - os.path.join(shared.cmd_opts.ckpt_dir, '**/*.vae.pt'), - os.path.join(shared.cmd_opts.ckpt_dir, '**/*.vae.safetensors'), - ] - if shared.cmd_opts.vae_dir is not None and os.path.isdir(shared.cmd_opts.vae_dir): paths += [ os.path.join(shared.cmd_opts.vae_dir, '**/*.ckpt'), diff --git a/modules/ui_extra_networks_checkpoints.py b/modules/ui_extra_networks_checkpoints.py index d69d144d..9a4ee153 100644 --- a/modules/ui_extra_networks_checkpoints.py +++ b/modules/ui_extra_networks_checkpoints.py @@ -45,7 +45,7 @@ class ExtraNetworksPageCheckpoints(ui_extra_networks.ExtraNetworksPage): yield item def allowed_directories_for_previews(self): - return [v for v in [shared.cmd_opts.ckpt_dir, sd_models.model_path] if v is not None] + return [v for v in (*shared.cmd_opts.ckpt_dirs, sd_models.model_path) if os.path.isdir(str(v))] def create_user_metadata_editor(self, ui, tabname): return CheckpointUserMetadataEditor(ui, tabname, self)