From aff56bab39e03958edb972aa512594f359d9b190 Mon Sep 17 00:00:00 2001 From: pkpk Date: Sun, 4 Feb 2024 12:35:58 +0900 Subject: [PATCH] Changed ckpt_type specification --- LoadTempCheckpoint.py | 5 +++-- LoadTempLoRA.py | 5 +++-- 2 files changed, 6 insertions(+), 4 deletions(-) diff --git a/LoadTempCheckpoint.py b/LoadTempCheckpoint.py index 0612268..8c7f183 100644 --- a/LoadTempCheckpoint.py +++ b/LoadTempCheckpoint.py @@ -18,7 +18,7 @@ class LoadTempCheckpoint: return { "required": { "ckpt_url": ("STRING", {"default": ""}), - "ckpt_type": (["safetensors", "other"], {"default": "safetensors"}), + "ckpt_type": (["auto", "safetensors", "other"], {"default": "auto"}), "download_split": ("INT", {"default": 4, "min": 1, "max": 8, "step": 1}) }, } @@ -34,7 +34,8 @@ class LoadTempCheckpoint: if bin is None: raise file_name if file_name is not None else Exception("Download failed.") - sd = load_torch_bin(bin, ckpt_type=="safetensors" or file_name.endswith(".safetensors")) + is_safetensors = file_name.endswith(".safetensors") if ckpt_type =="auto" else ckpt_type == "safetensors" + sd = load_torch_bin(bin, is_safetensors) sd_keys = sd.keys() clip = None clipvision = None diff --git a/LoadTempLoRA.py b/LoadTempLoRA.py index a1bbdd9..b26c0ab 100644 --- a/LoadTempLoRA.py +++ b/LoadTempLoRA.py @@ -19,7 +19,7 @@ class LoadTempLoRA: "model": ("MODEL",), "clip": ("CLIP", ), "ckpt_url": ("STRING", {"default": ""}), - "ckpt_type": (["safetensors", "other"], {"default": "safetensors"}), + "ckpt_type": (["auto", "safetensors", "other"], {"default": "auto"}), "download_split": ("INT", {"default": 4, "min": 1, "max": 8, "step": 1}), "strength_model": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}), "strength_clip": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}), @@ -48,7 +48,8 @@ class LoadTempLoRA: if bin is None: raise file_name if file_name is not None else Exception("Download failed.") - lora = load_torch_bin(bin, ckpt_type=="safetensors" or file_name.endswith(".safetensors"), safe_load=True) + is_safetensors = file_name.endswith(".safetensors") if ckpt_type =="auto" else ckpt_type == "safetensors" + lora = load_torch_bin(bin, is_safetensors, safe_load=True) self.loaded_lora = (ckpt_url, lora) model_lora, clip_lora = comfy.sd.load_lora_for_models(model, clip, lora, strength_model, strength_clip)