From 165f1320fd26a7a904e7fe558843891ef56fa802 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 15 Apr 2024 22:00:40 +0300 Subject: [PATCH] Add ipadapter loading from .safetensors --- ip_adapter/ip_adapter.py | 16 +++++++++++++++- 1 file changed, 15 insertions(+), 1 deletion(-) diff --git a/ip_adapter/ip_adapter.py b/ip_adapter/ip_adapter.py index 479f43b..99f87f8 100644 --- a/ip_adapter/ip_adapter.py +++ b/ip_adapter/ip_adapter.py @@ -31,8 +31,22 @@ class IPAdapter: self.dtype = dtype # load ip adapter model - ipadapter_model = torch.load(ipadapter_ckpt_path, map_location="cpu") + from comfy.utils import load_torch_file + ipadapter_model = load_torch_file(ipadapter_ckpt_path, safe_load=True) + if ipadapter_ckpt_path.lower().endswith(".safetensors"): + st_model = {"image_proj": {}, "ip_adapter": {}} + for key in ipadapter_model.keys(): + if key.startswith("image_proj."): + st_model["image_proj"][key.replace("image_proj.", "")] = ipadapter_model[key] + elif key.startswith("ip_adapter."): + st_model["ip_adapter"][key.replace("ip_adapter.", "")] = ipadapter_model[key] + ipadapter_model = st_model + del st_model + + if not "ip_adapter" in ipadapter_model.keys() or not ipadapter_model["ip_adapter"]: + raise Exception("invalid IPAdapter model {}".format(ipadapter_ckpt_path)) + # detect features self.is_plus = "latents" in ipadapter_model["image_proj"] self.output_cross_attention_dim = ipadapter_model["ip_adapter"]["1.to_k_ip.weight"].shape[1]