first commit

This commit is contained in:
pkpk
2024-02-02 11:04:05 +09:00
commit a577477b6d
6 changed files with 289 additions and 0 deletions
+95
View File
@@ -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)
+55
View File
@@ -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)
+16
View File
@@ -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.
+13
View File
@@ -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']
+1
View File
@@ -0,0 +1 @@
requests
+109
View File
@@ -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