commit a577477b6d743c839c95b21652644df13725c9ba Author: pkpk Date: Fri Feb 2 11:04:05 2024 +0900 first commit diff --git a/LoadTempCheckpoint.py b/LoadTempCheckpoint.py new file mode 100644 index 0000000..0612268 --- /dev/null +++ b/LoadTempCheckpoint.py @@ -0,0 +1,95 @@ +import torch + +import folder_paths + +import comfy.sd +import comfy.utils +import comfy.clip_vision +import comfy.model_detection +import comfy.model_management +import comfy.model_patcher +import comfy.checkpoint_pickle + +from .utils import download_file, load_torch_bin + +class LoadTempCheckpoint: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "ckpt_url": ("STRING", {"default": ""}), + "ckpt_type": (["safetensors", "other"], {"default": "safetensors"}), + "download_split": ("INT", {"default": 4, "min": 1, "max": 8, "step": 1}) + }, + } + RETURN_TYPES = ("MODEL", "CLIP", "VAE", "CLIP_VISION") + FUNCTION = "load_checkpoint" + + CATEGORY = "temporary_loaders" + + def load_checkpoint(self, ckpt_url, ckpt_type, download_split, output_model=True, output_vae=True, output_clip=True, output_clipvision=True): + + bin, file_name = download_file(ckpt_url, download_split) + + 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")) + sd_keys = sd.keys() + clip = None + clipvision = None + vae = None + model = None + model_patcher = None + clip_target = None + + parameters = comfy.utils.calculate_parameters(sd, "model.diffusion_model.") + unet_dtype = comfy.model_management.unet_dtype(model_params=parameters) + load_device = comfy.model_management.get_torch_device() + manual_cast_dtype = comfy.model_management.unet_manual_cast(unet_dtype, load_device) + + class WeightsLoader(torch.nn.Module): + pass + + model_config = comfy.model_detection.model_config_from_unet(sd, "model.diffusion_model.", unet_dtype) + model_config.set_manual_cast(manual_cast_dtype) + + if model_config is None: + raise RuntimeError("ERROR: Could not detect model type of: {}".format(ckpt_url)) + + if model_config.clip_vision_prefix is not None: + if output_clipvision: + clipvision = comfy.clip_vision.load_clipvision_from_sd(sd, model_config.clip_vision_prefix, True) + + if output_model: + inital_load_device = comfy.model_management.unet_inital_load_device(parameters, unet_dtype) + offload_device = comfy.model_management.unet_offload_device() + model = model_config.get_model(sd, "model.diffusion_model.", device=inital_load_device) + model.load_model_weights(sd, "model.diffusion_model.") + + if output_vae: + # vae_sd = comfy.utils.state_dict_prefix_replace(sd, {k: "" for k in model_config.vae_key_prefix}, filter_keys=True) + vae_sd = comfy.utils.state_dict_prefix_replace(sd, {k: "" for k in "first_stage_model."}, filter_keys=True) + vae_sd = model_config.process_vae_state_dict(vae_sd) + vae = comfy.sd.VAE(sd=vae_sd) + + if output_clip: + w = WeightsLoader() + clip_target = model_config.clip_target() + if clip_target is not None: + clip = comfy.sd.CLIP(clip_target, embedding_directory=folder_paths.get_folder_paths("embeddings")) + w.cond_stage_model = clip.cond_stage_model + sd = model_config.process_clip_state_dict(sd) + comfy.sd.load_model_weights(w, sd) + + left_over = sd.keys() + if len(left_over) > 0: + print("left over keys:", left_over) + + if output_model: + model_patcher = comfy.model_patcher.ModelPatcher(model, load_device=load_device, offload_device=comfy.model_management.unet_offload_device(), current_device=inital_load_device) + if inital_load_device != torch.device("cpu"): + print("loaded straight to GPU") + comfy.model_management.load_model_gpu(model_patcher) + + return (model_patcher, clip, vae, clipvision) diff --git a/LoadTempLoRA.py b/LoadTempLoRA.py new file mode 100644 index 0000000..a1bbdd9 --- /dev/null +++ b/LoadTempLoRA.py @@ -0,0 +1,55 @@ +import comfy.sd +import comfy.utils +import comfy.clip_vision +import comfy.model_detection +import comfy.model_management +import comfy.model_patcher +import comfy.checkpoint_pickle + +from .utils import download_file, load_torch_bin + +class LoadTempLoRA: + def __init__(self): + self.loaded_lora = None + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "clip": ("CLIP", ), + "ckpt_url": ("STRING", {"default": ""}), + "ckpt_type": (["safetensors", "other"], {"default": "safetensors"}), + "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}), + }, + } + RETURN_TYPES = ("MODEL", "CLIP") + FUNCTION = "load_lora" + + CATEGORY = "temporary_loaders" + + def load_lora(self, model, clip, ckpt_url, ckpt_type, download_split, strength_model, strength_clip): + if strength_model == 0 and strength_clip == 0: + return (model, clip) + + lora = None + if self.loaded_lora is not None: + if self.loaded_lora[0] == ckpt_url: + lora = self.loaded_lora[1] + else: + temp = self.loaded_lora + self.loaded_lora = None + del temp + + if lora is None: + bin, file_name = download_file(ckpt_url, download_split) + 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) + self.loaded_lora = (ckpt_url, lora) + + model_lora, clip_lora = comfy.sd.load_lora_for_models(model, clip, lora, strength_model, strength_clip) + return (model_lora, clip_lora) diff --git a/README.md b/README.md new file mode 100644 index 0000000..d69205d --- /dev/null +++ b/README.md @@ -0,0 +1,16 @@ +# ComfyUI-TemporaryLoader +This is a custom node of ComfyUI that downloads and loads models from the input URL. The model is temporarily downloaded into memory and not saved to storage. + +This could be useful when trying out models or when using various models on machines with limited storage. Since the model is downloaded into memory, expect higher memory usage than usual. + +## Installation: +1. Use `git clone https://github.com/pkpkTech/ComfyUI-TemporaryLoader` in your ComfyUI custom nodes directory +1. Use `pip install -r requirements.txt` in ComfyUI-ngrok directory + +## Usage +The "Load Checkpoint (Temporary)" and "Load LoRA (Temporary)" nodes will be added to the temporary_loader category. + +In addition to the standard Load node, the following items are added +- ckpt_url: URL of the model +- ckpt_type: Select whether the model you want to use is safetensors or other. If the file name extension of the downloaded file is `.safetensors`, it will be loaded as safetensors even if you select other. +- download_split: Specify the number of splits for parallel downloading. diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..463e914 --- /dev/null +++ b/__init__.py @@ -0,0 +1,13 @@ +from .LoadTempCheckpoint import LoadTempCheckpoint +from .LoadTempLoRA import LoadTempLoRA + +NODE_CLASS_MAPPINGS = { + "LoadTempCheckpoint": LoadTempCheckpoint, + "LoadTempLoRA": LoadTempLoRA +} +NODE_DISPLAY_NAME_MAPPINGS = { + "LoadTempCheckpoint": "Load Checkpoint (Temporary)", + "LoadTempLoRA": "Load LoRA (Temporary)" +} + +__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..663bd1f --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +requests \ No newline at end of file diff --git a/utils.py b/utils.py new file mode 100644 index 0000000..b2f8290 --- /dev/null +++ b/utils.py @@ -0,0 +1,109 @@ +import requests +import threading +import io +from urllib.parse import unquote +from tqdm.auto import tqdm + +import safetensors.torch +import torch + +import comfy.utils + +def download_chunk(url, start_byte, end_byte, result_parts, total_size, pbar_web, pbar_cli): + thr = total_size // 10 + size = 0 + cnt = 0 + headers = {"Range": f"bytes={start_byte}-{end_byte}"} + with requests.get(url, headers=headers, stream=True, allow_redirects=True) as response: + response.raise_for_status() + with io.BytesIO() as part_data: + for chunk in response.iter_content(chunk_size=1024): + if chunk: + part_data.write(chunk) + chunk_size = len(chunk) + size+=chunk_size + pbar_cli.update(chunk_size) + if size >= thr: + size = 0 + if cnt < 9: + cnt+=1 + pbar_web.update(1) + + result_parts[start_byte] = part_data.getvalue() + del part_data + pbar_web.update(10 - cnt) + +def download_file(url, num_threads=4): + try: + response = requests.get(url, stream=True, allow_redirects=True) + response.raise_for_status() + total_size = int(response.headers.get("content-length", 0)) + content_disposition = response.headers.get("Content-Disposition", None) + except requests.exceptions.RequestException as e: + return None, e + finally: + response.close() + + chunk_size = total_size // num_threads + + pbar = comfy.utils.ProgressBar(num_threads * 10 + 1) + pbar.update_absolute(0) + + file_name = "blank" + if content_disposition: + # Content-Dispositionヘッダーからファイル名を抽出する + parts = content_disposition.split(";") + for part in parts: + if part.strip().startswith("filename="): + file_name = unquote(part.strip().split("=")[1].strip('"')) + break + + threads = [] + result_parts = {} + with tqdm(total=total_size) as pbar_cli: + for i in range(num_threads): + start_byte = chunk_size * i + end_byte = start_byte + chunk_size - 1 if i < num_threads - 1 else "" + chunk_total_size = total_size - start_byte if i == num_threads-1 else chunk_size + thread = threading.Thread(target=download_chunk, args=(url, start_byte, end_byte, result_parts, chunk_total_size, pbar, pbar_cli)) + thread.start() + threads.append(thread) + + for thread in threads: + thread.join() + + pbar.update_absolute(num_threads * 10 + 1) + + # 全てのチャンクがダウンロードできているか確認 + if len(result_parts) < num_threads: + return None, None + + # ダウンロードされた部分データを結合してバイナリデータとして返す + sorted_parts = sorted(result_parts.items()) + result = b"".join(part for start_byte, part in sorted_parts) + + return result, file_name + +# Bin to torch +def load_torch_bin(bin, is_safetensors, safe_load=False, device=None): + if device is None: + device = torch.device("cpu") + if is_safetensors: + sd = safetensors.torch.load(bin) + else: + if safe_load: + if not "weights_only" in torch.load.__code__.co_varnames: + print("Warning torch.load doesn't support weights_only on this pytorch version, loading unsafely.") + safe_load = False + ckpt = io.BytesIO(bin) + if safe_load: + pl_sd = torch.load(ckpt, map_location=device, weights_only=True) + else: + pl_sd = torch.load(ckpt, map_location=device, pickle_module=comfy.checkpoint_pickle) + if "global_step" in pl_sd: + print(f"Global Step: {pl_sd['global_step']}") + if "state_dict" in pl_sd: + sd = pl_sd["state_dict"] + else: + sd = pl_sd + return sd