Updated to new IPAdapter + fixes

This commit is contained in:
IDGallagher
2024-05-13 17:22:15 +01:00
parent 5e207a8c7c
commit 2a4768b32f
8 changed files with 2467 additions and 471 deletions
+2 -2
View File
@@ -565,11 +565,11 @@ class BatchCreativeInterpolationNode:
for i, bin in enumerate(bins):
ipadapter_application = IPAdapterBatchImport()
negative_noise = torch.cat(bin.noiseBatch, dim=0) if len(bin.noiseBatch) > 0 else None
model, = ipadapter_application.apply_ipadapter(model=model, ipadapter=ipadapter, image=torch.cat(bin.imageBatch, dim=0), weight=[x * base_ipa_advanced_settings["ipa_weight"] for x in bin.weight_schedule], weight_type=base_ipa_advanced_settings["ipa_weight_type"], start_at=base_ipa_advanced_settings["ipa_starts_at"], end_at=base_ipa_advanced_settings["ipa_ends_at"], clip_vision=clip_vision,image_negative=negative_noise,embeds_scaling=base_ipa_advanced_settings["ipa_embeds_scaling"], image_schedule=bin.image_schedule)
model, *_ = ipadapter_application.apply_ipadapter(model=model, ipadapter=ipadapter, image=torch.cat(bin.imageBatch, dim=0), weight=[x * base_ipa_advanced_settings["ipa_weight"] for x in bin.weight_schedule], weight_type=base_ipa_advanced_settings["ipa_weight_type"], start_at=base_ipa_advanced_settings["ipa_starts_at"], end_at=base_ipa_advanced_settings["ipa_ends_at"], clip_vision=clip_vision,image_negative=negative_noise,embeds_scaling=base_ipa_advanced_settings["ipa_embeds_scaling"], encode_batch_size=1, image_schedule=bin.image_schedule)
if high_detail_mode:
tiled_ipa_application = IPAdapterTiledBatchImport()
negative_noise = torch.cat(bin.bigNoiseBatch, dim=0) if len(bin.bigNoiseBatch) > 0 else None
model, *_ = tiled_ipa_application.apply_tiled(model=model, ipadapter=ipadapter, image=torch.cat(bin.bigImageBatch, dim=0), weight=[x * detail_ipa_advanced_settings["ipa_weight"] for x in bin.weight_schedule], weight_type=detail_ipa_advanced_settings["ipa_weight_type"], start_at=detail_ipa_advanced_settings["ipa_starts_at"], end_at=detail_ipa_advanced_settings["ipa_ends_at"], clip_vision=clip_vision,sharpening=0.1,image_negative=negative_noise,embeds_scaling=detail_ipa_advanced_settings["ipa_embeds_scaling"], image_schedule=bin.image_schedule)
model, *_ = tiled_ipa_application.apply_tiled(model=model, ipadapter=ipadapter, image=torch.cat(bin.bigImageBatch, dim=0), weight=[x * detail_ipa_advanced_settings["ipa_weight"] for x in bin.weight_schedule], weight_type=detail_ipa_advanced_settings["ipa_weight_type"], start_at=detail_ipa_advanced_settings["ipa_starts_at"], end_at=detail_ipa_advanced_settings["ipa_ends_at"], clip_vision=clip_vision,sharpening=0.1,image_negative=negative_noise,embeds_scaling=detail_ipa_advanced_settings["ipa_embeds_scaling"], encode_batch_size=1, image_schedule=bin.image_schedule)
comparison_diagram, = plot_weight_comparison(all_cn_frame_numbers, all_cn_weights, all_ipa_frame_numbers, all_ipa_weights, buffer)
return comparison_diagram, positive, negative, model, sparse_indexes, last_key_frame_position, buffer
@@ -4,9 +4,202 @@ import torch.nn.functional as F
from comfy.ldm.modules.attention import optimized_attention
from .utils import tensor_to_size
class CrossAttentionPatchImport:
class Attn2ReplaceImport:
def __init__(self, callback=None, **kwargs):
self.callback = [callback]
self.kwargs = [kwargs]
def add(self, callback, **kwargs):
self.callback.append(callback)
self.kwargs.append(kwargs)
for key, value in kwargs.items():
setattr(self, key, value)
def __call__(self, q, k, v, extra_options):
dtype = q.dtype
out = optimized_attention(q, k, v, extra_options["n_heads"])
sigma = extra_options["sigmas"].detach().cpu()[0].item() if 'sigmas' in extra_options else 999999999.9
for i, callback in enumerate(self.callback):
if sigma <= self.kwargs[i]["sigma_start"] and sigma >= self.kwargs[i]["sigma_end"]:
out = out + callback(out, q, k, v, extra_options, **self.kwargs[i])
return out.to(dtype=dtype)
def ipadapter_attention_import(out, q, k, v, extra_options, module_key='', ipadapter=None, weight=1.0, cond=None, cond_alt=None, uncond=None, weight_type="linear", mask=None, sigma_start=0.0, sigma_end=1.0, unfold_batch=False, image_schedule=None, embeds_scaling='V only', **kwargs):
dtype = q.dtype
cond_or_uncond = extra_options["cond_or_uncond"]
block_type = extra_options["block"][0]
#block_id = extra_options["block"][1]
t_idx = extra_options["transformer_index"]
layers = 11 if '101_to_k_ip' in ipadapter.ip_layers.to_kvs else 16
k_key = module_key + "_to_k_ip"
v_key = module_key + "_to_v_ip"
# extra options for AnimateDiff
ad_params = extra_options['ad_params'] if "ad_params" in extra_options else None
b = q.shape[0]
seq_len = q.shape[1]
batch_prompt = b // len(cond_or_uncond)
_, _, oh, ow = extra_options["original_shape"]
if weight_type == 'ease in':
weight = weight * (0.05 + 0.95 * (1 - t_idx / layers))
elif weight_type == 'ease out':
weight = weight * (0.05 + 0.95 * (t_idx / layers))
elif weight_type == 'ease in-out':
weight = weight * (0.05 + 0.95 * (1 - abs(t_idx - (layers/2)) / (layers/2)))
elif weight_type == 'reverse in-out':
weight = weight * (0.05 + 0.95 * (abs(t_idx - (layers/2)) / (layers/2)))
elif weight_type == 'weak input' and block_type == 'input':
weight = weight * 0.2
elif weight_type == 'weak middle' and block_type == 'middle':
weight = weight * 0.2
elif weight_type == 'weak output' and block_type == 'output':
weight = weight * 0.2
elif weight_type == 'strong middle' and (block_type == 'input' or block_type == 'output'):
weight = weight * 0.2
elif isinstance(weight, dict):
if t_idx not in weight:
return 0
weight = weight[t_idx]
if cond_alt is not None and t_idx in cond_alt:
cond = cond_alt[t_idx]
del cond_alt
if unfold_batch:
# Check AnimateDiff context window
if ad_params is not None and ad_params["sub_idxs"] is not None:
if isinstance(weight, torch.Tensor):
weight = tensor_to_size(weight, ad_params["full_length"])
weight = torch.Tensor(weight[ad_params["sub_idxs"]])
if torch.all(weight == 0):
return 0
weight = weight.repeat(len(cond_or_uncond), 1, 1) # repeat for cond and uncond
elif weight == 0:
return 0
if image_schedule is not None:
# Use the image_schedule as a lookup table to get the embedded image corresponding to each sub_idx
# If image_schedule isn't long enough then use the last image
cond_idxs = [image_schedule[i if i < len(image_schedule) else -1] for i in ad_params["sub_idxs"]]
cond = torch.Tensor(cond[cond_idxs])
uncond = torch.Tensor(uncond[cond_idxs])
else:
# if image length matches or exceeds full_length get sub_idx images
if cond.shape[0] >= ad_params["full_length"]:
cond = torch.Tensor(cond[ad_params["sub_idxs"]])
uncond = torch.Tensor(uncond[ad_params["sub_idxs"]])
# otherwise get sub_idxs images
else:
cond = tensor_to_size(cond, ad_params["full_length"])
uncond = tensor_to_size(uncond, ad_params["full_length"])
cond = cond[ad_params["sub_idxs"]]
uncond = uncond[ad_params["sub_idxs"]]
else:
if isinstance(weight, torch.Tensor):
weight = tensor_to_size(weight, batch_prompt)
if torch.all(weight == 0):
return 0
weight = weight.repeat(len(cond_or_uncond), 1, 1) # repeat for cond and uncond
elif weight == 0:
return 0
cond = tensor_to_size(cond, batch_prompt)
uncond = tensor_to_size(uncond, batch_prompt)
k_cond = ipadapter.ip_layers.to_kvs[k_key](cond)
k_uncond = ipadapter.ip_layers.to_kvs[k_key](uncond)
v_cond = ipadapter.ip_layers.to_kvs[v_key](cond)
v_uncond = ipadapter.ip_layers.to_kvs[v_key](uncond)
else:
# TODO: should we always convert the weights to a tensor?
if isinstance(weight, torch.Tensor):
weight = tensor_to_size(weight, batch_prompt)
if torch.all(weight == 0):
return 0
weight = weight.repeat(len(cond_or_uncond), 1, 1) # repeat for cond and uncond
elif weight == 0:
return 0
k_cond = ipadapter.ip_layers.to_kvs[k_key](cond).repeat(batch_prompt, 1, 1)
k_uncond = ipadapter.ip_layers.to_kvs[k_key](uncond).repeat(batch_prompt, 1, 1)
v_cond = ipadapter.ip_layers.to_kvs[v_key](cond).repeat(batch_prompt, 1, 1)
v_uncond = ipadapter.ip_layers.to_kvs[v_key](uncond).repeat(batch_prompt, 1, 1)
ip_k = torch.cat([(k_cond, k_uncond)[i] for i in cond_or_uncond], dim=0)
ip_v = torch.cat([(v_cond, v_uncond)[i] for i in cond_or_uncond], dim=0)
if embeds_scaling == 'K+mean(V) w/ C penalty':
scaling = float(ip_k.shape[2]) / 1280.0
weight = weight * scaling
ip_k = ip_k * weight
ip_v_mean = torch.mean(ip_v, dim=1, keepdim=True)
ip_v = (ip_v - ip_v_mean) + ip_v_mean * weight
out_ip = optimized_attention(q, ip_k, ip_v, extra_options["n_heads"])
del ip_v_mean
elif embeds_scaling == 'K+V w/ C penalty':
scaling = float(ip_k.shape[2]) / 1280.0
weight = weight * scaling
ip_k = ip_k * weight
ip_v = ip_v * weight
out_ip = optimized_attention(q, ip_k, ip_v, extra_options["n_heads"])
elif embeds_scaling == 'K+V':
ip_k = ip_k * weight
ip_v = ip_v * weight
out_ip = optimized_attention(q, ip_k, ip_v, extra_options["n_heads"])
else:
#ip_v = ip_v * weight
out_ip = optimized_attention(q, ip_k, ip_v, extra_options["n_heads"])
out_ip = out_ip * weight # I'm doing this to get the same results as before
if mask is not None:
mask_h = oh / math.sqrt(oh * ow / seq_len)
mask_h = int(mask_h) + int((seq_len % int(mask_h)) != 0)
mask_w = seq_len // mask_h
# check if using AnimateDiff and sliding context window
if (mask.shape[0] > 1 and ad_params is not None and ad_params["sub_idxs"] is not None):
# if mask length matches or exceeds full_length, get sub_idx masks
if mask.shape[0] >= ad_params["full_length"]:
mask = torch.Tensor(mask[ad_params["sub_idxs"]])
mask = F.interpolate(mask.unsqueeze(1), size=(mask_h, mask_w), mode="bilinear").squeeze(1)
else:
mask = F.interpolate(mask.unsqueeze(1), size=(mask_h, mask_w), mode="bilinear").squeeze(1)
mask = tensor_to_size(mask, ad_params["full_length"])
mask = mask[ad_params["sub_idxs"]]
else:
mask = F.interpolate(mask.unsqueeze(1), size=(mask_h, mask_w), mode="bilinear").squeeze(1)
mask = tensor_to_size(mask, batch_prompt)
mask = mask.repeat(len(cond_or_uncond), 1, 1)
mask = mask.view(mask.shape[0], -1, 1).repeat(1, 1, out.shape[2])
# covers cases where extreme aspect ratios can cause the mask to have a wrong size
mask_len = mask_h * mask_w
if mask_len < seq_len:
pad_len = seq_len - mask_len
pad1 = pad_len // 2
pad2 = pad_len - pad1
mask = F.pad(mask, (0, 0, pad1, pad2), value=0.0)
elif mask_len > seq_len:
crop_start = (mask_len - seq_len) // 2
mask = mask[:, crop_start:crop_start+seq_len, :]
out_ip = out_ip * mask
#out = out + out_ip
return out_ip.to(dtype=dtype)
"""
class CrossAttentionPatch:
# forward for patching
def __init__(self, ipadapter=None, number=0, weight=1.0, cond=None, cond_alt=None, uncond=None, weight_type="linear", mask=None, sigma_start=0.0, sigma_end=1.0, unfold_batch=False, image_schedule=None, embeds_scaling='V only'):
def __init__(self, ipadapter=None, number=0, weight=1.0, cond=None, cond_alt=None, uncond=None, weight_type="linear", mask=None, sigma_start=0.0, sigma_end=1.0, unfold_batch=False, embeds_scaling='V only'):
self.weights = [weight]
self.ipadapters = [ipadapter]
self.conds = [cond]
@@ -17,7 +210,6 @@ class CrossAttentionPatchImport:
self.sigma_starts = [sigma_start]
self.sigma_ends = [sigma_end]
self.unfold_batch = [unfold_batch]
self.image_schedule = [image_schedule]
self.embeds_scaling = [embeds_scaling]
self.number = number
self.layers = 11 if '101_to_k_ip' in ipadapter.ip_layers.to_kvs else 16 # TODO: check if this is a valid condition to detect all models
@@ -25,27 +217,7 @@ class CrossAttentionPatchImport:
self.k_key = str(self.number*2+1) + "_to_k_ip"
self.v_key = str(self.number*2+1) + "_to_v_ip"
@classmethod
def from_cross_attention_patch(cls, patch):
instance = cls(ipadapter = patch.ipadapters[0])
instance.weights = patch.weights
instance.ipadapters = patch.ipadapters
instance.conds = patch.conds
instance.conds_alt = patch.conds_alt
instance.unconds = patch.unconds
instance.weight_types = patch.weight_types
instance.masks = patch.masks
instance.sigma_starts = patch.sigma_starts
instance.sigma_ends = patch.sigma_ends
instance.unfold_batch = patch.unfold_batch
instance.embeds_scaling = patch.embeds_scaling
instance.number = patch.number
instance.layers = patch.layers
instance.k_key = patch.k_key
instance.v_key = patch.v_key
return instance
def set_new_condition(self, ipadapter=None, number=0, weight=1.0, cond=None, cond_alt=None, uncond=None, weight_type="linear", mask=None, sigma_start=0.0, sigma_end=1.0, unfold_batch=False, image_schedule=None, embeds_scaling='V only'):
def set_new_condition(self, ipadapter=None, number=0, weight=1.0, cond=None, cond_alt=None, uncond=None, weight_type="linear", mask=None, sigma_start=0.0, sigma_end=1.0, unfold_batch=False, embeds_scaling='V only'):
self.weights.append(weight)
self.ipadapters.append(ipadapter)
self.conds.append(cond)
@@ -56,7 +228,6 @@ class CrossAttentionPatchImport:
self.sigma_starts.append(sigma_start)
self.sigma_ends.append(sigma_end)
self.unfold_batch.append(unfold_batch)
self.image_schedule.append(image_schedule)
self.embeds_scaling.append(embeds_scaling)
def __call__(self, q, k, v, extra_options):
@@ -76,7 +247,7 @@ class CrossAttentionPatchImport:
out = optimized_attention(q, k, v, extra_options["n_heads"])
_, _, oh, ow = extra_options["original_shape"]
for weight, cond, cond_alt, uncond, ipadapter, mask, weight_type, sigma_start, sigma_end, unfold_batch, image_schedule, embeds_scaling in zip(self.weights, self.conds, self.conds_alt, self.unconds, self.ipadapters, self.masks, self.weight_types, self.sigma_starts, self.sigma_ends, self.unfold_batch, self.image_schedule, self.embeds_scaling):
for weight, cond, cond_alt, uncond, ipadapter, mask, weight_type, sigma_start, sigma_end, unfold_batch, embeds_scaling in zip(self.weights, self.conds, self.conds_alt, self.unconds, self.ipadapters, self.masks, self.weight_types, self.sigma_starts, self.sigma_ends, self.unfold_batch, self.embeds_scaling):
if sigma <= sigma_start and sigma >= sigma_end:
if weight_type == 'ease in':
weight = weight * (0.05 + 0.95 * (1 - t_idx / self.layers))
@@ -116,23 +287,16 @@ class CrossAttentionPatchImport:
elif weight == 0:
continue
if image_schedule is not None:
# Use the image_schedule as a lookup table to get the embedded image corresponding to each sub_idx
# If image_schedule isn't long enough then use the last image
cond_idxs = [image_schedule[i if i < len(image_schedule) else -1] for i in ad_params["sub_idxs"]]
cond = torch.Tensor(cond[cond_idxs])
uncond = torch.Tensor(uncond[cond_idxs])
# if image length matches or exceeds full_length get sub_idx images
if cond.shape[0] >= ad_params["full_length"]:
cond = torch.Tensor(cond[ad_params["sub_idxs"]])
uncond = torch.Tensor(uncond[ad_params["sub_idxs"]])
# otherwise get sub_idxs images
else:
# if image length matches or exceeds full_length get sub_idx images
if cond.shape[0] >= ad_params["full_length"]:
cond = torch.Tensor(cond[ad_params["sub_idxs"]])
uncond = torch.Tensor(uncond[ad_params["sub_idxs"]])
# otherwise get sub_idxs images
else:
cond = tensor_to_size(cond, ad_params["full_length"])
uncond = tensor_to_size(uncond, ad_params["full_length"])
cond = cond[ad_params["sub_idxs"]]
uncond = uncond[ad_params["sub_idxs"]]
cond = tensor_to_size(cond, ad_params["full_length"])
uncond = tensor_to_size(uncond, ad_params["full_length"])
cond = cond[ad_params["sub_idxs"]]
uncond = uncond[ad_params["sub_idxs"]]
else:
if isinstance(weight, torch.Tensor):
weight = tensor_to_size(weight, batch_prompt)
@@ -228,3 +392,4 @@ class CrossAttentionPatchImport:
out = out + out_ip
return out.to(dtype=dtype)
"""
+367 -71
View File
@@ -4,6 +4,7 @@ import math
import folder_paths
import comfy.model_management as model_management
from node_helpers import conditioning_set_values
from comfy.clip_vision import load as load_clip_vision
from comfy.sd import load_lora_for_models
import comfy.utils
@@ -16,7 +17,7 @@ except ImportError:
import torchvision.transforms as T
from .image_proj_models import MLPProjModelImport, MLPProjModelFaceIdImport, ProjModelFaceIdPlusImport, ResamplerImport, ImageProjModelImport
from .CrossAttentionPatchImport import CrossAttentionPatchImport
from .CrossAttentionPatchImport import Attn2ReplaceImport, ipadapter_attention_import
from .utils import (
encode_image_masked,
tensor_to_size,
@@ -134,19 +135,22 @@ class To_KV(nn.Module):
self.to_kvs[key.replace(".weight", "").replace(".", "_")].weight.data = value
def set_model_patch_replace(model, patch_kwargs, key):
to = model.model_options["transformer_options"]
to = model.model_options["transformer_options"].copy()
if "patches_replace" not in to:
to["patches_replace"] = {}
else:
to["patches_replace"] = to["patches_replace"].copy()
if "attn2" not in to["patches_replace"]:
to["patches_replace"]["attn2"] = {}
if key not in to["patches_replace"]["attn2"]:
to["patches_replace"]["attn2"][key] = CrossAttentionPatchImport(**patch_kwargs)
else:
if isinstance(to["patches_replace"]["attn2"][key], CrossAttentionPatchImport):
to["patches_replace"]["attn2"][key].set_new_condition(**patch_kwargs)
else:
to["patches_replace"]["attn2"][key] = CrossAttentionPatchImport.from_cross_attention_patch(to["patches_replace"]["attn2"][key])
to["patches_replace"]["attn2"][key].set_new_condition(**patch_kwargs)
to["patches_replace"]["attn2"] = to["patches_replace"]["attn2"].copy()
if key not in to["patches_replace"]["attn2"]:
to["patches_replace"]["attn2"][key] = Attn2ReplaceImport(ipadapter_attention_import, **patch_kwargs)
model.model_options["transformer_options"] = to
else:
to["patches_replace"]["attn2"][key].add(ipadapter_attention_import, **patch_kwargs)
def ipadapter_execute(model,
ipadapter,
@@ -168,11 +172,12 @@ def ipadapter_execute(model,
unfold_batch=False,
image_schedule=None,
embeds_scaling='V only',
layer_weights=None):
layer_weights=None,
encode_batch_size=0,):
device = model_management.get_torch_device()
dtype = model_management.unet_dtype()
if dtype not in [torch.float32, torch.float16, torch.bfloat16]:
dtype = torch.float16 if comfy.model_management.should_use_fp16() else torch.float32
dtype = torch.float16 if model_management.should_use_fp16() else torch.float32
is_full = "proj.3.weight" in ipadapter["image_proj"]
is_portrait = "proj.2.weight" in ipadapter["image_proj"] and not "proj.3.weight" in ipadapter["image_proj"] and not "0.to_q_lora.down.weight" in ipadapter["ip_adapter"]
@@ -196,7 +201,7 @@ def ipadapter_execute(model,
print("\033[33mINFO: the IPAdapter reference image is not a square, CLIPImageProcessor will resize and crop it at the center. If the main focus of the picture is not in the middle the result might not be what you are expecting.\033[0m")
if isinstance(weight, list):
weight = torch.tensor(weight).unsqueeze(-1).unsqueeze(-1).to(device, dtype=dtype) if unfold_batch else weight[0]
weight = torch.tensor(weight).unsqueeze(-1).unsqueeze(-1).to(device, dtype=dtype) if unfold_batch else weight[0]
# special weight types
if layer_weights is not None and layer_weights != '':
@@ -244,7 +249,7 @@ def ipadapter_execute(model,
face_cond_embeds.append(torch.from_numpy(face[0].normed_embedding).unsqueeze(0))
else:
face_cond_embeds.append(torch.from_numpy(face[0].embedding).unsqueeze(0))
image.append(image_to_tensor(face_align.norm_crop(image_iface[i], landmark=face[0].kps, image_size=256)))
image.append(image_to_tensor(face_align.norm_crop(image_iface[i], landmark=face[0].kps, image_size=256 if is_sdxl else 224)))
if 640 not in size:
print(f"\033[33mINFO: InsightFace detection resolution lowered to {size}.\033[0m")
@@ -256,25 +261,27 @@ def ipadapter_execute(model,
del image_iface, face
if image is not None:
img_cond_embeds = encode_image_masked(clipvision, image)
img_cond_embeds = encode_image_masked(clipvision, image, batch_size=encode_batch_size)
if image_composition is not None:
img_comp_cond_embeds = encode_image_masked(clipvision, image_composition)
img_comp_cond_embeds = encode_image_masked(clipvision, image_composition, batch_size=encode_batch_size)
if is_plus:
img_cond_embeds = img_cond_embeds.penultimate_hidden_states
image_negative = image_negative if image_negative is not None else torch.zeros([1, 224, 224, 3])
img_uncond_embeds = encode_image_masked(clipvision, image_negative).penultimate_hidden_states
img_uncond_embeds = encode_image_masked(clipvision, image_negative, batch_size=encode_batch_size).penultimate_hidden_states
if image_composition is not None:
img_comp_cond_embeds = img_comp_cond_embeds.penultimate_hidden_states
else:
img_cond_embeds = img_cond_embeds.image_embeds if not is_faceid else face_cond_embeds
if image_negative is not None and not is_faceid:
img_uncond_embeds = encode_image_masked(clipvision, image_negative).image_embeds
img_uncond_embeds = encode_image_masked(clipvision, image_negative, batch_size=encode_batch_size).image_embeds
else:
img_uncond_embeds = torch.zeros_like(img_cond_embeds)
if image_composition is not None:
img_comp_cond_embeds = img_comp_cond_embeds.image_embeds
del image, image_negative, image_composition
del image_negative, image_composition
image = None if not is_faceid else image # if it's face_id we need the cropped face for later
elif pos_embed is not None:
img_cond_embeds = pos_embed
@@ -366,7 +373,6 @@ def ipadapter_execute(model,
patch_kwargs = {
"ipadapter": ipa,
"number": 0,
"weight": weight,
"cond": cond,
"cond_alt": cond_alt,
@@ -380,31 +386,37 @@ def ipadapter_execute(model,
"embeds_scaling": embeds_scaling,
}
number = 0
if not is_sdxl:
for id in [1,2,4,5,7,8]: # id of input_blocks that have cross attention
patch_kwargs["module_key"] = str(number*2+1)
set_model_patch_replace(model, patch_kwargs, ("input", id))
patch_kwargs["number"] += 1
number += 1
for id in [3,4,5,6,7,8,9,10,11]: # id of output_blocks that have cross attention
patch_kwargs["module_key"] = str(number*2+1)
set_model_patch_replace(model, patch_kwargs, ("output", id))
patch_kwargs["number"] += 1
number += 1
patch_kwargs["module_key"] = str(number*2+1)
set_model_patch_replace(model, patch_kwargs, ("middle", 0))
else:
for id in [4,5,7,8]: # id of input_blocks that have cross attention
block_indices = range(2) if id in [4, 5] else range(10) # transformer_depth
for index in block_indices:
patch_kwargs["module_key"] = str(number*2+1)
set_model_patch_replace(model, patch_kwargs, ("input", id, index))
patch_kwargs["number"] += 1
number += 1
for id in range(6): # id of output_blocks that have cross attention
block_indices = range(2) if id in [3, 4, 5] else range(10) # transformer_depth
for index in block_indices:
patch_kwargs["module_key"] = str(number*2+1)
set_model_patch_replace(model, patch_kwargs, ("output", id, index))
patch_kwargs["number"] += 1
number += 1
for index in range(10):
patch_kwargs["module_key"] = str(number*2+1)
set_model_patch_replace(model, patch_kwargs, ("middle", 0, index))
patch_kwargs["number"] += 1
return model
number += 1
return (model, image)
"""
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
@@ -607,7 +619,7 @@ class IPAdapterSimple:
if 'clipvision' not in ipadapter:
raise Exception("CLIPVision model not present in the pipeline. Please load the models with the IPAdapterUnifiedLoader node.")
return (ipadapter_execute(model.clone(), ipadapter['ipadapter']['model'], ipadapter['clipvision']['model'], **ipa_args), )
return ipadapter_execute(model.clone(), ipadapter['ipadapter']['model'], ipadapter['clipvision']['model'], **ipa_args)
class IPAdapterAdvancedImport:
def __init__(self):
@@ -638,9 +650,18 @@ class IPAdapterAdvancedImport:
FUNCTION = "apply_ipadapter"
CATEGORY = "ipadapter"
def apply_ipadapter(self, model, ipadapter, start_at, end_at, weight = 1.0, weight_style=1.0, weight_composition=1.0, expand_style=False, weight_type="linear", combine_embeds="concat", weight_faceidv2=None, image=None, image_style=None, image_composition=None, image_negative=None, clip_vision=None, image_schedule=None, attn_mask=None, insightface=None, embeds_scaling='V only', layer_weights=None):
def apply_ipadapter(self, model, ipadapter, start_at=0.0, end_at=1.0, weight=1.0, weight_style=1.0, weight_composition=1.0, expand_style=False, weight_type="linear", combine_embeds="concat", weight_faceidv2=None, image=None, image_style=None, image_composition=None, image_negative=None, clip_vision=None, attn_mask=None, insightface=None, embeds_scaling='V only', layer_weights=None, ipadapter_params=None, encode_batch_size=0, image_schedule=None):
is_sdxl = isinstance(model.model, (comfy.model_base.SDXL, comfy.model_base.SDXLRefiner, comfy.model_base.SDXL_instructpix2pix))
if 'ipadapter' in ipadapter:
ipadapter_model = ipadapter['ipadapter']['model']
clip_vision = clip_vision if clip_vision is not None else ipadapter['clipvision']['model']
else:
ipadapter_model = ipadapter
if clip_vision is None:
raise Exception("Missing CLIPVision model.")
if image_style is not None: # we are doing style + composition transfer
if not is_sdxl:
raise Exception("Style + Composition transfer is only available for SDXL models at the moment.") # TODO: check feasibility for SD1.5 models
@@ -651,39 +672,49 @@ class IPAdapterAdvancedImport:
image_composition = image_style
weight_type = "strong style and composition" if expand_style else "style and composition"
ipa_args = {
"image": image,
"image_composition": image_composition,
"image_negative": image_negative,
"weight": weight,
"weight_composition": weight_composition,
"weight_faceidv2": weight_faceidv2,
"weight_type": weight_type,
"combine_embeds": combine_embeds,
"start_at": start_at,
"end_at": end_at,
"attn_mask": attn_mask,
"unfold_batch": self.unfold_batch,
"image_schedule": image_schedule,
"embeds_scaling": embeds_scaling,
"insightface": insightface if insightface is not None else ipadapter['insightface']['model'] if 'insightface' in ipadapter else None,
"layer_weights": layer_weights,
}
if 'ipadapter' in ipadapter:
ipadapter_model = ipadapter['ipadapter']['model']
clip_vision = clip_vision if clip_vision is not None else ipadapter['clipvision']['model']
if ipadapter_params is not None: # we are doing batch processing
image = ipadapter_params['image']
attn_mask = ipadapter_params['attn_mask']
weight = ipadapter_params['weight']
weight_type = ipadapter_params['weight_type']
start_at = ipadapter_params['start_at']
end_at = ipadapter_params['end_at']
else:
ipadapter_model = ipadapter
clip_vision = clip_vision
# at this point weight can be a list from the batch-weight or a single float
weight = [weight]
if clip_vision is None:
raise Exception("Missing CLIPVision model.")
image = image if isinstance(image, list) else [image]
work_model = model.clone()
for i in range(len(image)):
if image[i] is None:
continue
ipa_args = {
"image": image[i],
"image_composition": image_composition,
"image_negative": image_negative,
"weight": weight[i],
"weight_composition": weight_composition,
"weight_faceidv2": weight_faceidv2,
"weight_type": weight_type if not isinstance(weight_type, list) else weight_type[i],
"combine_embeds": combine_embeds,
"start_at": start_at if not isinstance(start_at, list) else start_at[i],
"end_at": end_at if not isinstance(end_at, list) else end_at[i],
"attn_mask": attn_mask if not isinstance(attn_mask, list) else attn_mask[i],
"unfold_batch": self.unfold_batch,
"image_schedule": image_schedule,
"embeds_scaling": embeds_scaling,
"insightface": insightface if insightface is not None else ipadapter['insightface']['model'] if 'insightface' in ipadapter else None,
"layer_weights": layer_weights,
"encode_batch_size": encode_batch_size,
}
work_model, face_image = ipadapter_execute(work_model, ipadapter_model, clip_vision, **ipa_args)
del ipadapter
return (ipadapter_execute(model.clone(), ipadapter_model, clip_vision, **ipa_args), )
return (work_model, face_image, )
class IPAdapterBatchImport(IPAdapterAdvancedImport):
def __init__(self):
@@ -701,6 +732,7 @@ class IPAdapterBatchImport(IPAdapterAdvancedImport):
"start_at": ("FLOAT", { "default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001 }),
"end_at": ("FLOAT", { "default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001 }),
"embeds_scaling": (['V only', 'K+V', 'K+V w/ C penalty', 'K+mean(V) w/ C penalty'], ),
"encode_batch_size": ("INT", { "default": 0, "min": 0, "max": 4096 }),
},
"optional": {
"image_negative": ("IMAGE",),
@@ -788,6 +820,8 @@ class IPAdapterFaceID(IPAdapterAdvancedImport):
}
CATEGORY = "ipadapter/faceid"
RETURN_TYPES = ("MODEL","IMAGE",)
RETURN_NAMES = ("MODEL", "face_image", )
class IPAAdapterFaceIDBatch(IPAdapterFaceID):
def __init__(self):
@@ -824,7 +858,7 @@ class IPAdapterTiledImport:
FUNCTION = "apply_tiled"
CATEGORY = "ipadapter/tiled"
def apply_tiled(self, model, ipadapter, image, weight, weight_type, start_at, end_at, sharpening, combine_embeds="concat", image_negative=None, attn_mask=None, clip_vision=None, embeds_scaling='V only', image_schedule=None):
def apply_tiled(self, model, ipadapter, image, weight, weight_type, start_at, end_at, sharpening, image_schedule=None, combine_embeds="concat", image_negative=None, attn_mask=None, clip_vision=None, embeds_scaling='V only', encode_batch_size=0):
# 1. Select the models
if 'ipadapter' in ipadapter:
ipadapter_model = ipadapter['ipadapter']['model']
@@ -922,11 +956,12 @@ class IPAdapterTiledImport:
"attn_mask": masks[i],
"unfold_batch": self.unfold_batch,
"embeds_scaling": embeds_scaling,
"encode_batch_size": encode_batch_size,
"image_schedule": image_schedule,
}
# apply the ipadapter to the model without cloning it
model = ipadapter_execute(model, ipadapter_model, clip_vision, **ipa_args)
model, _ = ipadapter_execute(model, ipadapter_model, clip_vision, **ipa_args)
return (model, torch.cat(tiles), torch.cat(masks), )
@@ -947,6 +982,7 @@ class IPAdapterTiledBatchImport(IPAdapterTiledImport):
"end_at": ("FLOAT", { "default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001 }),
"sharpening": ("FLOAT", { "default": 0.0, "min": 0.0, "max": 1.0, "step": 0.05 }),
"embeds_scaling": (['V only', 'K+V', 'K+V w/ C penalty', 'K+mean(V) w/ C penalty'], ),
"encode_batch_size": ("INT", { "default": 0, "min": 0, "max": 4096 }),
},
"optional": {
"image_negative": ("IMAGE",),
@@ -1005,7 +1041,7 @@ class IPAdapterEmbeds:
del ipadapter
return (ipadapter_execute(model.clone(), ipadapter_model, clip_vision, **ipa_args), )
return ipadapter_execute(model.clone(), ipadapter_model, clip_vision, **ipa_args)
class IPAdapterMS(IPAdapterAdvancedImport):
@classmethod
@@ -1034,6 +1070,25 @@ class IPAdapterMS(IPAdapterAdvancedImport):
CATEGORY = "ipadapter/dev"
class IPAdapterFromParams(IPAdapterAdvancedImport):
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL", ),
"ipadapter": ("IPADAPTER", ),
"ipadapter_params": ("IPADAPTER_PARAMS", ),
"combine_embeds": (["concat", "add", "subtract", "average", "norm average"],),
"embeds_scaling": (['V only', 'K+V', 'K+V w/ C penalty', 'K+mean(V) w/ C penalty'], ),
},
"optional": {
"image_negative": ("IMAGE",),
"clip_vision": ("CLIP_VISION",),
}
}
CATEGORY = "ipadapter/params"
"""
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
Helpers
@@ -1336,27 +1391,57 @@ class IPAdapterWeights:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"weights": ("STRING", {"default": '1.0', "multiline": True }),
"timing": (["custom", "linear", "ease_in_out", "ease_in", "ease_out", "reverse_in_out", "random"], ),
"weights": ("STRING", {"default": '1.0, 0.0', "multiline": True }),
"timing": (["custom", "linear", "ease_in_out", "ease_in", "ease_out", "random"], { "default": "linear" } ),
"frames": ("INT", {"default": 0, "min": 0, "max": 9999, "step": 1 }),
"start_frame": ("INT", {"default": 0, "min": 0, "max": 9999, "step": 1 }),
"end_frame": ("INT", {"default": 9999, "min": 0, "max": 9999, "step": 1 }),
},
"add_starting_frames": ("INT", {"default": 0, "min": 0, "max": 9999, "step": 1 }),
"add_ending_frames": ("INT", {"default": 0, "min": 0, "max": 9999, "step": 1 }),
"method": (["full batch", "shift batches", "alternate batches"], { "default": "full batch" }),
}, "optional": {
"image": ("IMAGE",),
}
}
RETURN_TYPES = ("FLOAT",)
RETURN_TYPES = ("FLOAT", "FLOAT", "INT", "IMAGE", "IMAGE", "WEIGHTS_STRATEGY")
RETURN_NAMES = ("weights", "weights_invert", "total_frames", "image_1", "image_2", "weights_strategy")
FUNCTION = "weights"
CATEGORY = "ipadapter/weights"
CATEGORY = "ipadapter/utils"
def weights(self, weights, timing, frames, start_frame, end_frame):
def weights(self, weights='', timing='custom', frames=0, start_frame=0, end_frame=9999, add_starting_frames=0, add_ending_frames=0, method='full batch', weights_strategy=None, image=None):
import random
frame_count = image.shape[0] if image is not None else 0
if weights_strategy is not None:
weights = weights_strategy["weights"]
timing = weights_strategy["timing"]
frames = weights_strategy["frames"]
start_frame = weights_strategy["start_frame"]
end_frame = weights_strategy["end_frame"]
add_starting_frames = weights_strategy["add_starting_frames"]
add_ending_frames = weights_strategy["add_ending_frames"]
method = weights_strategy["method"]
frame_count = weights_strategy["frame_count"]
else:
weights_strategy = {
"weights": weights,
"timing": timing,
"frames": frames,
"start_frame": start_frame,
"end_frame": end_frame,
"add_starting_frames": add_starting_frames,
"add_ending_frames": add_ending_frames,
"method": method,
"frame_count": frame_count,
}
# convert the string to a list of floats separated by commas or newlines
weights = weights.replace("\n", ",")
weights = [float(weight) for weight in weights.split(",") if weight.strip() != ""]
if timing != "custom":
frames = max(frames, 2)
start = 0.0
end = 1.0
@@ -1381,16 +1466,227 @@ class IPAdapterWeights:
weights.append(start + (end - start) * math.sin(i / n * math.pi / 2))
elif timing == "ease_out":
weights.append(start + (end - start) * (1 - math.cos(i / n * math.pi / 2)))
elif timing == "reverse_in_out":
weights.append(start + (end - start) * (1 - math.sin((1 - i / n) * math.pi / 2)))
elif timing == "random":
weights.append(random.uniform(start, end))
weights[-1] = end if timing != "random" else weights[-1]
weights[-1] = end if timing != "random" else weights[-1]
if end_frame < frames:
weights.extend([end] * (frames - end_frame))
if len(weights) == 0:
weights = [0.0]
return (weights, )
frames = len(weights)
# repeat the images for cross fade
image_1 = None
image_2 = None
if image is not None:
if "shift" in method:
image_1 = image[:-1]
image_2 = image[1:]
weights = weights * image_1.shape[0]
image_1 = image_1.repeat_interleave(frames, 0)
image_2 = image_2.repeat_interleave(frames, 0)
elif "alternate" in method:
image_1 = image[::2].repeat_interleave(2, 0)
image_1 = image_1[1:]
image_2 = image[1::2].repeat_interleave(2, 0)
mew_weights = weights + [1.0 - w for w in weights]
mew_weights = mew_weights * (image_1.shape[0] // 2)
if image.shape[0] % 2:
image_1 = image_1[:-1]
else:
image_2 = image_2[:-1]
mew_weights = mew_weights + weights
weights = mew_weights
image_1 = image_1.repeat_interleave(frames, 0)
image_2 = image_2.repeat_interleave(frames, 0)
else:
weights = weights * image.shape[0]
image_1 = image.repeat_interleave(frames, 0)
# add starting and ending frames
if add_starting_frames > 0:
weights = [weights[0]] * add_starting_frames + weights
image_1 = torch.cat([image[:1].repeat(add_starting_frames, 1, 1, 1), image_1], dim=0)
if image_2 is not None:
image_2 = torch.cat([image[:1].repeat(add_starting_frames, 1, 1, 1), image_2], dim=0)
if add_ending_frames > 0:
weights = weights + [weights[-1]] * add_ending_frames
image_1 = torch.cat([image_1, image[-1:].repeat(add_ending_frames, 1, 1, 1)], dim=0)
if image_2 is not None:
image_2 = torch.cat([image_2, image[-1:].repeat(add_ending_frames, 1, 1, 1)], dim=0)
weights_invert = [1.0 - w for w in weights]
frame_count = len(weights)
return (weights, weights_invert, frame_count, image_1, image_2, weights_strategy,)
class IPAdapterWeightsFromStrategy(IPAdapterWeights):
@classmethod
def INPUT_TYPES(s):
return {"required": {
"weights_strategy": ("WEIGHTS_STRATEGY",),
}, "optional": {
"image": ("IMAGE",),
}
}
class IPAdapterPromptScheduleFromWeightsStrategy():
@classmethod
def INPUT_TYPES(s):
return {"required": {
"weights_strategy": ("WEIGHTS_STRATEGY",),
"prompt": ("STRING", {"default": "", "multiline": True }),
}}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("prompt_schedule", )
FUNCTION = "prompt_schedule"
CATEGORY = "ipadapter/weights"
def prompt_schedule(self, weights_strategy, prompt=""):
frames = weights_strategy["frames"]
add_starting_frames = weights_strategy["add_starting_frames"]
add_ending_frames = weights_strategy["add_ending_frames"]
frame_count = weights_strategy["frame_count"]
out = ""
prompt = [p for p in prompt.split("\n") if p.strip() != ""]
if len(prompt) > 0 and frame_count > 0:
# prompt_pos must be the same size as the image batch
if len(prompt) > frame_count:
prompt = prompt[:frame_count]
elif len(prompt) < frame_count:
prompt += [prompt[-1]] * (frame_count - len(prompt))
if add_starting_frames > 0:
out += f"\"0\": \"{prompt[0]}\",\n"
for i in range(frame_count):
out += f"\"{i * frames + add_starting_frames}\": \"{prompt[i]}\",\n"
if add_ending_frames > 0:
out += f"\"{frame_count * frames + add_starting_frames}\": \"{prompt[-1]}\",\n"
return (out, )
class IPAdapterCombineWeights:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"weights_1": ("FLOAT", { "default": 0.0, "min": 0.0, "max": 1.0, "step": 0.05 }),
"weights_2": ("FLOAT", { "default": 0.0, "min": 0.0, "max": 1.0, "step": 0.05 }),
}}
RETURN_TYPES = ("FLOAT", "INT")
RETURN_NAMES = ("weights", "count")
FUNCTION = "combine"
CATEGORY = "ipadapter/utils"
def combine(self, weights_1, weights_2):
if not isinstance(weights_1, list):
weights_1 = [weights_1]
if not isinstance(weights_2, list):
weights_2 = [weights_2]
weights = weights_1 + weights_2
return (weights, len(weights), )
class IPAdapterRegionalConditioning:
@classmethod
def INPUT_TYPES(s):
return {"required": {
#"set_cond_area": (["default", "mask bounds"],),
"image": ("IMAGE",),
"image_weight": ("FLOAT", { "default": 1.0, "min": -1.0, "max": 3.0, "step": 0.05 }),
"prompt_weight": ("FLOAT", { "default": 1.0, "min": 0.0, "max": 10.0, "step": 0.05 }),
"weight_type": (WEIGHT_TYPES, ),
"start_at": ("FLOAT", { "default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001 }),
"end_at": ("FLOAT", { "default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001 }),
}, "optional": {
"mask": ("MASK",),
"positive": ("CONDITIONING",),
"negative": ("CONDITIONING",),
}}
RETURN_TYPES = ("IPADAPTER_PARAMS", "CONDITIONING", "CONDITIONING", )
RETURN_NAMES = ("IPADAPTER_PARAMS", "POSITIVE", "NEGATIVE")
FUNCTION = "conditioning"
CATEGORY = "ipadapter/params"
def conditioning(self, image, image_weight, prompt_weight, weight_type, start_at, end_at, mask=None, positive=None, negative=None):
set_area_to_bounds = False #if set_cond_area == "default" else True
if mask is not None:
if positive is not None:
positive = conditioning_set_values(positive, {"mask": mask, "set_area_to_bounds": set_area_to_bounds, "mask_strength": prompt_weight})
if negative is not None:
negative = conditioning_set_values(negative, {"mask": mask, "set_area_to_bounds": set_area_to_bounds, "mask_strength": prompt_weight})
ipadapter_params = {
"image": [image],
"attn_mask": [mask],
"weight": [image_weight],
"weight_type": [weight_type],
"start_at": [start_at],
"end_at": [end_at],
}
return (ipadapter_params, positive, negative, )
class IPAdapterCombineParams:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"params_1": ("IPADAPTER_PARAMS",),
"params_2": ("IPADAPTER_PARAMS",),
}, "optional": {
"params_3": ("IPADAPTER_PARAMS",),
"params_4": ("IPADAPTER_PARAMS",),
"params_5": ("IPADAPTER_PARAMS",),
}}
RETURN_TYPES = ("IPADAPTER_PARAMS",)
FUNCTION = "combine"
CATEGORY = "ipadapter/params"
def combine(self, params_1, params_2, params_3=None, params_4=None, params_5=None):
ipadapter_params = {
"image": params_1["image"] + params_2["image"],
"attn_mask": params_1["attn_mask"] + params_2["attn_mask"],
"weight": params_1["weight"] + params_2["weight"],
"weight_type": params_1["weight_type"] + params_2["weight_type"],
"start_at": params_1["start_at"] + params_2["start_at"],
"end_at": params_1["end_at"] + params_2["end_at"],
}
if params_3 is not None:
ipadapter_params["image"] += params_3["image"]
ipadapter_params["attn_mask"] += params_3["attn_mask"]
ipadapter_params["weight"] += params_3["weight"]
ipadapter_params["weight_type"] += params_3["weight_type"]
ipadapter_params["start_at"] += params_3["start_at"]
ipadapter_params["end_at"] += params_3["end_at"]
if params_4 is not None:
ipadapter_params["image"] += params_4["image"]
ipadapter_params["attn_mask"] += params_4["attn_mask"]
ipadapter_params["weight"] += params_4["weight"]
ipadapter_params["weight_type"] += params_4["weight_type"]
ipadapter_params["start_at"] += params_4["start_at"]
ipadapter_params["end_at"] += params_4["end_at"]
if params_5 is not None:
ipadapter_params["image"] += params_5["image"]
ipadapter_params["attn_mask"] += params_5["attn_mask"]
ipadapter_params["weight"] += params_5["weight"]
ipadapter_params["weight_type"] += params_5["weight_type"]
ipadapter_params["start_at"] += params_5["start_at"]
ipadapter_params["end_at"] += params_5["end_at"]
return (ipadapter_params, )
+23 -8
View File
@@ -27,6 +27,14 @@ Please consider a [Github Sponsorship](https://github.com/sponsors/cubiq) or [Pa
## Important updates
**2024/05/02**: Add `encode_batch_size` to the Advanced batch node. This can be useful for animations with a lot of frames to reduce the VRAM usage during the image encoding. Please note that results will be slightly different based on the batch size.
**2024/04/27**: Refactored the IPAdapterWeights mostly useful for AnimateDiff animations.
**2024/04/21**: Added Regional Conditioning nodes to simplify attention masking and masked text conditioning.
**2024/04/16**: Added support for the new SDXL portrait unnorm model (link below). It's very strong and tends to ignore the text conditioning. Lower the CFG to 3-4 or use a RescaleCFG node.
**2024/04/12**: Added scheduled weights. Useful for animations.
**2024/04/09**: Added experimental Style/Composition transfer for SD1.5. The results are often not as good as SDXL. Optimal weight seems to be from 0.8 to 2.0. The **Style+Composition node doesn't work for SD1.5** at the moment, you can only alter either the Style or the Composition, I need more time for testing. Old workflows will still work **but you may need to refresh the page and re-select the weight type!**
@@ -37,9 +45,9 @@ Please consider a [Github Sponsorship](https://github.com/sponsors/cubiq) or [Pa
**2024/03/27**: Added Style transfer weight type for SDXL
**2024/03/23**: Complete code rewrite!. **This is a breaking update!** Your previous workflows won't work and you'll need to recreate them. You've been warned! After the update, refresh your browser, delete the old IPAdapter nodes and create the new ones.
**2024/03/23**: Complete code rewrite! **This is a breaking update!** Your previous workflows won't work and you'll need to recreate them. You've been warned! After the update, refresh your browser, delete the old IPAdapter nodes and create the new ones.
*(I removed all previous updates because they were about the previous version of the extension)*
*(I removed old updates related to the previous version of the extension)*
## Example workflows
@@ -53,11 +61,12 @@ The [examples directory](./examples/) has many workflows that cover all IPAdapte
<img src="https://img.youtube.com/vi/_JzDcgKgghY/hqdefault.jpg" alt="Watch the video" />
</a>
**:star: [New IPAdapter features](https://youtu.be/_JzDcgKgghY)**
- **:star: [New IPAdapter features](https://youtu.be/_JzDcgKgghY)**
- **:art: [IPAdapter Style and Composition](https://www.youtube.com/watch?v=czcgJnoDVd4)**
The following videos are about the previous version of IPAdapter, but they still contain valuable information.
**:nerd_face: [Basic usage video](https://youtu.be/7m9ZZFU3HWo)**, **:rocket: [Advanced features video](https://www.youtube.com/watch?v=mJQ62ly7jrg)**, **:japanese_goblin: [Attention Masking video](https://www.youtube.com/watch?v=vqG1VXKteQg)**, **:movie_camera: [Animation Features video](https://www.youtube.com/watch?v=ddYbhv3WgWw)**
:nerd_face: [Basic usage video](https://youtu.be/7m9ZZFU3HWo), :rocket: [Advanced features video](https://www.youtube.com/watch?v=mJQ62ly7jrg), :japanese_goblin: [Attention Masking video](https://www.youtube.com/watch?v=vqG1VXKteQg), :movie_camera: [Animation Features video](https://www.youtube.com/watch?v=ddYbhv3WgWw)
## Installation
@@ -85,7 +94,7 @@ Remember you can also use any custom location setting an `ipadapter` entry in th
**FaceID** models require `insightface`, you need to install it in your ComfyUI environment. Check [this issue](https://github.com/cubiq/ComfyUI_IPAdapter_plus/issues/162) for help. Remember that most FaceID models also need a LoRA.
For the Unified Loader to work the files need to be named exactly as shown in the table below.
For the Unified Loader to work the files need to be named exactly as shown in the list below.
- `/ComfyUI/models/ipadapter`
- [ip-adapter-faceid_sd15.bin](https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid_sd15.bin), base FaceID model
@@ -94,6 +103,7 @@ For the Unified Loader to work the files need to be named exactly as shown in th
- [ip-adapter-faceid_sdxl.bin](https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid_sdxl.bin), SDXL base FaceID
- [ip-adapter-faceid-plusv2_sdxl.bin](https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-plusv2_sdxl.bin), SDXL plus v2
- [ip-adapter-faceid-portrait_sdxl.bin](https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-portrait_sdxl.bin), SDXL text prompt style transfer
- [ip-adapter-faceid-portrait_sdxl_unnorm.bin](https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-portrait_sdxl_unnorm.bin), very strong style transfer SDXL only
- **Deprecated** [ip-adapter-faceid-plus_sd15.bin](https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-plus_sd15.bin), FaceID plus v1
- **Deprecated** [ip-adapter-faceid-portrait_sd15.bin](https://huggingface.co/h94/IP-Adapter-FaceID/resolve/main/ip-adapter-faceid-portrait_sd15.bin), v1 of the portrait model
@@ -136,9 +146,13 @@ Please check the [troubleshooting](https://github.com/cubiq/ComfyUI_IPAdapter_pl
It's only thanks to generous sponsors that **the whole community** can enjoy open and free software. Please join me in thanking the following companies and individuals!
### Gold sponsors
### :trophy: Gold sponsors
[![Kaiber.ai](https://f.latent.vision/imgs/kaiber.png)](https://kaiber.ai/)
[![Kaiber.ai](https://f.latent.vision/imgs/kaiber.png)](https://kaiber.ai/)&nbsp; &nbsp;[![Kaiber.ai](https://f.latent.vision/imgs/replicate.png)](https://replicate.com/)
### :tada: Silver sponsors
[![OperArt.ai](https://f.latent.vision/imgs/openart.png?r=1)](https://openart.ai/workflows)
### Companies supporting my projects
@@ -148,8 +162,9 @@ It's only thanks to generous sponsors that **the whole community** can enjoy ope
- [Jack Gane](https://github.com/ganeJackS)
- [Nathan Shipley](https://www.nathanshipley.com/)
- [Dkdnzia](https://github.com/Dkdnzia)
### One-time Extraordinaire
### One-time Extraordinaires
- [Eric Rollei](https://github.com/EricRollei)
- [francaleu](https://github.com/francaleu)
@@ -1,6 +1,6 @@
{
"last_node_id": 23,
"last_link_id": 43,
"last_link_id": 44,
"nodes": [
{
"id": 8,
@@ -14,7 +14,7 @@
"1": 46
},
"flags": {},
"order": 11,
"order": 10,
"mode": 0,
"inputs": [
{
@@ -87,7 +87,7 @@
"1": 262
},
"flags": {},
"order": 10,
"order": 9,
"mode": 0,
"inputs": [
{
@@ -146,7 +146,7 @@
"1": 582.3048095703125
},
"flags": {},
"order": 12,
"order": 11,
"mode": 0,
"inputs": [
{
@@ -211,7 +211,7 @@
"1": 180.6060791015625
},
"flags": {},
"order": 6,
"order": 5,
"mode": 0,
"inputs": [
{
@@ -249,7 +249,7 @@
"1": 164.31304931640625
},
"flags": {},
"order": 5,
"order": 4,
"mode": 0,
"inputs": [
{
@@ -335,7 +335,7 @@
"1": 126
},
"flags": {},
"order": 4,
"order": 3,
"mode": 0,
"inputs": [
{
@@ -391,7 +391,7 @@
"1": 78
},
"flags": {},
"order": 8,
"order": 7,
"mode": 0,
"inputs": [
{
@@ -440,10 +440,10 @@
],
"size": {
"0": 315,
"1": 166
"1": 190
},
"flags": {},
"order": 9,
"order": 8,
"mode": 0,
"inputs": [
{
@@ -460,7 +460,7 @@
{
"name": "image",
"type": "IMAGE",
"link": 41
"link": 44
},
{
"name": "attn_mask",
@@ -485,7 +485,8 @@
"widgets_values": [
0.4,
0,
1
1,
"standard"
]
},
{
@@ -497,10 +498,10 @@
],
"size": {
"0": 315,
"1": 298
"1": 322
},
"flags": {},
"order": 7,
"order": 6,
"mode": 0,
"inputs": [
{
@@ -549,6 +550,15 @@
],
"shape": 3,
"slot_index": 0
},
{
"name": "face_image",
"type": "IMAGE",
"links": [
44
],
"shape": 3,
"slot_index": 1
}
],
"properties": {
@@ -560,46 +570,8 @@
"linear",
"concat",
0,
1
]
},
{
"id": 23,
"type": "LoadImage",
"pos": [
1280,
-230
],
"size": {
"0": 315,
"1": 314
},
"flags": {},
"order": 3,
"mode": 0,
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
41
],
"shape": 3,
"slot_index": 0
},
{
"name": "MASK",
"type": "MASK",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"rosario.png",
"image"
1,
"V only"
]
}
],
@@ -724,14 +696,6 @@
0,
"MODEL"
],
[
41,
23,
0,
21,
2,
"IMAGE"
],
[
42,
21,
@@ -747,6 +711,14 @@
22,
0,
"MODEL"
],
[
44,
18,
1,
21,
2,
"IMAGE"
]
],
"groups": [],
File diff suppressed because it is too large Load Diff
@@ -1,6 +1,6 @@
{
"last_node_id": 21,
"last_link_id": 37,
"last_node_id": 22,
"last_link_id": 40,
"nodes": [
{
"id": 7,
@@ -14,7 +14,7 @@
"1": 180.6060791015625
},
"flags": {},
"order": 8,
"order": 5,
"mode": 0,
"inputs": [
{
@@ -80,49 +80,6 @@
"Node name for S&R": "VAEDecode"
}
},
{
"id": 5,
"type": "EmptyLatentImage",
"pos": [
801,
1097
],
"size": [
315,
106
],
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "batch_size",
"type": "INT",
"link": 35,
"widget": {
"name": "batch_size"
}
}
],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
2
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "EmptyLatentImage"
},
"widgets_values": [
512,
512,
6
]
},
{
"id": 6,
"type": "CLIPTextEncode",
@@ -135,7 +92,7 @@
"1": 164.31304931640625
},
"flags": {},
"order": 7,
"order": 4,
"mode": 0,
"inputs": [
{
@@ -220,170 +177,6 @@
1
]
},
{
"id": 12,
"type": "LoadImage",
"pos": [
311,
270
],
"size": {
"0": 315,
"1": 314
},
"flags": {},
"order": 0,
"mode": 0,
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
25
],
"shape": 3,
"slot_index": 0
},
{
"name": "MASK",
"type": "MASK",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"warrior_woman.png",
"image"
]
},
{
"id": 17,
"type": "PrepImageForClipVision",
"pos": [
797,
87
],
"size": {
"0": 315,
"1": 106
},
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 25
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
30
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "PrepImageForClipVision"
},
"widgets_values": [
"LANCZOS",
"top",
0.15
]
},
{
"id": 20,
"type": "IPAdapterWeights",
"pos": [
757,
318
],
"size": [
263.5047280787487,
183.75987616018006
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "frames",
"type": "INT",
"link": 34,
"widget": {
"name": "frames"
},
"slot_index": 0
}
],
"outputs": [
{
"name": "FLOAT",
"type": "FLOAT",
"links": [
32
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "IPAdapterWeights"
},
"widgets_values": [
"1.0,0.0",
"linear",
6,
0,
9999
]
},
{
"id": 21,
"type": "PrimitiveNode",
"pos": [
340,
1093
],
"size": {
"0": 210,
"1": 82
},
"flags": {},
"order": 1,
"mode": 0,
"outputs": [
{
"name": "INT",
"type": "INT",
"links": [
34,
35
],
"widget": {
"name": "frames"
},
"slot_index": 0
}
],
"title": "frames",
"properties": {
"Run widget replace on values": false
},
"widgets_values": [
6,
"fixed"
]
},
{
"id": 19,
"type": "IPAdapterBatch",
@@ -391,10 +184,10 @@
1173,
251
],
"size": [
315,
254
],
"size": {
"0": 315,
"1": 254
},
"flags": {},
"order": 9,
"mode": 0,
@@ -432,7 +225,7 @@
{
"name": "weight",
"type": "FLOAT",
"link": 32,
"link": 38,
"widget": {
"name": "weight"
},
@@ -473,7 +266,7 @@
"1": 78
},
"flags": {},
"order": 6,
"order": 3,
"mode": 0,
"inputs": [
{
@@ -525,7 +318,7 @@
"1": 98
},
"flags": {},
"order": 2,
"order": 0,
"mode": 0,
"outputs": [
{
@@ -561,6 +354,47 @@
"sd15/realisticVisionV51_v51VAE.safetensors"
]
},
{
"id": 17,
"type": "PrepImageForClipVision",
"pos": [
788,
43
],
"size": {
"0": 315,
"1": 106
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 25
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
30
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "PrepImageForClipVision"
},
"widgets_values": [
"LANCZOS",
"top",
0.15
]
},
{
"id": 9,
"type": "SaveImage",
@@ -568,10 +402,10 @@
1770,
710
],
"size": [
556.2374508110479,
892.1895739499892
],
"size": {
"0": 556.2374267578125,
"1": 892.1895751953125
},
"flags": {},
"order": 12,
"mode": 0,
@@ -586,6 +420,203 @@
"widgets_values": [
"IPAdapter"
]
},
{
"id": 12,
"type": "LoadImage",
"pos": [
311,
270
],
"size": {
"0": 315,
"1": 314
},
"flags": {},
"order": 1,
"mode": 0,
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
25
],
"shape": 3,
"slot_index": 0
},
{
"name": "MASK",
"type": "MASK",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"warrior_woman.png",
"image"
]
},
{
"id": 5,
"type": "EmptyLatentImage",
"pos": [
801,
1097
],
"size": [
309.1109879864148,
82
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "batch_size",
"type": "INT",
"link": 35,
"widget": {
"name": "batch_size"
}
}
],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
2
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "EmptyLatentImage"
},
"widgets_values": [
512,
512,
6
]
},
{
"id": 21,
"type": "PrimitiveNode",
"pos": [
340,
1093
],
"size": {
"0": 210,
"1": 82
},
"flags": {},
"order": 2,
"mode": 0,
"outputs": [
{
"name": "INT",
"type": "INT",
"links": [
35,
40
],
"widget": {
"name": "batch_size"
},
"slot_index": 0
}
],
"title": "frames",
"properties": {
"Run widget replace on values": false
},
"widgets_values": [
6,
"fixed"
]
},
{
"id": 22,
"type": "IPAdapterWeights",
"pos": [
761,
208
],
"size": [
299.9049990375719,
324.00000762939453
],
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": null
},
{
"name": "frames",
"type": "INT",
"link": 40,
"widget": {
"name": "frames"
}
}
],
"outputs": [
{
"name": "weights",
"type": "FLOAT",
"links": [
38
],
"shape": 3,
"slot_index": 0
},
{
"name": "weights_invert",
"type": "FLOAT",
"links": null,
"shape": 3
},
{
"name": "total_frames",
"type": "INT",
"links": null,
"shape": 3
},
{
"name": "image_1",
"type": "IMAGE",
"links": null,
"shape": 3
},
{
"name": "image_2",
"type": "IMAGE",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "IPAdapterWeights"
},
"widgets_values": [
"1.0, 0.0",
"linear",
6,
0,
9999,
0,
0,
"full batch"
]
}
],
"links": [
@@ -685,22 +716,6 @@
0,
"MODEL"
],
[
32,
20,
0,
19,
6,
"FLOAT"
],
[
34,
21,
0,
20,
0,
"INT"
],
[
35,
21,
@@ -724,6 +739,22 @@
19,
0,
"MODEL"
],
[
38,
22,
0,
19,
6,
"FLOAT"
],
[
40,
21,
0,
22,
1,
"INT"
]
],
"groups": [],
+53 -48
View File
@@ -15,9 +15,9 @@ def get_clipvision_file(preset):
clipvision_list = folder_paths.get_filename_list("clip_vision")
if preset.startswith("vit-g"):
pattern = '(ViT.bigG.14.*39B.b160k|ipadapter.*sdxl|sdxl.*model\.(bin|safetensors))'
pattern = r'(ViT.bigG.14.*39B.b160k|ipadapter.*sdxl|sdxl.*model\.(bin|safetensors))'
else:
pattern = '(ViT.H.14.*s32B.b79K|ipadapter.*sd15|sd1.?5.*model\.(bin|safetensors))'
pattern = r'(ViT.H.14.*s32B.b79K|ipadapter.*sd15|sd1.?5.*model\.(bin|safetensors))'
clipvision_file = [e for e in clipvision_list if re.search(pattern, e, re.IGNORECASE)]
clipvision_file = folder_paths.get_full_path("clip_vision", clipvision_file[0]) if clipvision_file else None
@@ -33,77 +33,77 @@ def get_ipadapter_file(preset, is_sdxl):
if preset.startswith("light"):
if is_sdxl:
raise Exception("light model is not supported for SDXL")
pattern = 'sd15.light.v11\.(safetensors|bin)$'
pattern = r'sd15.light.v11\.(safetensors|bin)$'
# if v11 is not found, try with the old version
if not [e for e in ipadapter_list if re.search(pattern, e, re.IGNORECASE)]:
pattern = 'sd15.light\.(safetensors|bin)$'
pattern = r'sd15.light\.(safetensors|bin)$'
elif preset.startswith("standard"):
if is_sdxl:
pattern = 'ip.adapter.sdxl.vit.h\.(safetensors|bin)$'
pattern = r'ip.adapter.sdxl.vit.h\.(safetensors|bin)$'
else:
pattern = 'ip.adapter.sd15\.(safetensors|bin)$'
pattern = r'ip.adapter.sd15\.(safetensors|bin)$'
elif preset.startswith("vit-g"):
if is_sdxl:
pattern = 'ip.adapter.sdxl\.(safetensors|bin)$'
pattern = r'ip.adapter.sdxl\.(safetensors|bin)$'
else:
pattern = 'sd15.vit.g\.(safetensors|bin)$'
pattern = r'sd15.vit.g\.(safetensors|bin)$'
elif preset.startswith("plus ("):
if is_sdxl:
pattern = 'plus.sdxl.vit.h\.(safetensors|bin)$'
pattern = r'plus.sdxl.vit.h\.(safetensors|bin)$'
else:
pattern = 'ip.adapter.plus.sd15\.(safetensors|bin)$'
pattern = r'ip.adapter.plus.sd15\.(safetensors|bin)$'
elif preset.startswith("plus face"):
if is_sdxl:
pattern = 'plus.face.sdxl.vit.h\.(safetensors|bin)$'
pattern = r'plus.face.sdxl.vit.h\.(safetensors|bin)$'
else:
pattern = 'plus.face.sd15\.(safetensors|bin)$'
pattern = r'plus.face.sd15\.(safetensors|bin)$'
elif preset.startswith("full"):
if is_sdxl:
raise Exception("full face model is not supported for SDXL")
pattern = 'full.face.sd15\.(safetensors|bin)$'
pattern = r'full.face.sd15\.(safetensors|bin)$'
elif preset.startswith("faceid portrait ("):
if is_sdxl:
pattern = 'portrait.sdxl\.(safetensors|bin)$'
pattern = r'portrait.sdxl\.(safetensors|bin)$'
else:
pattern = 'portrait.v11.sd15\.(safetensors|bin)$'
pattern = r'portrait.v11.sd15\.(safetensors|bin)$'
# if v11 is not found, try with the old version
if not [e for e in ipadapter_list if re.search(pattern, e, re.IGNORECASE)]:
pattern = 'portrait.sd15\.(safetensors|bin)$'
pattern = r'portrait.sd15\.(safetensors|bin)$'
is_insightface = True
elif preset.startswith("faceid portrait unnorm"):
if is_sdxl:
pattern = 'portrait.sdxl.unnorm\.(safetensors|bin)$'
pattern = r'portrait.sdxl.unnorm\.(safetensors|bin)$'
else:
raise Exception("portrait unnorm model is not supported for SD1.5")
is_insightface = True
elif preset == "faceid":
if is_sdxl:
pattern = 'faceid.sdxl\.(safetensors|bin)$'
lora_pattern = 'faceid.sdxl.lora\.safetensors$'
pattern = r'faceid.sdxl\.(safetensors|bin)$'
lora_pattern = r'faceid.sdxl.lora\.safetensors$'
else:
pattern = 'faceid.sd15\.(safetensors|bin)$'
lora_pattern = 'faceid.sd15.lora\.safetensors$'
pattern = r'faceid.sd15\.(safetensors|bin)$'
lora_pattern = r'faceid.sd15.lora\.safetensors$'
is_insightface = True
elif preset.startswith("faceid plus -"):
if is_sdxl:
raise Exception("faceid plus model is not supported for SDXL")
pattern = 'faceid.plus.sd15\.(safetensors|bin)$'
lora_pattern = 'faceid.plus.sd15.lora\.safetensors$'
pattern = r'faceid.plus.sd15\.(safetensors|bin)$'
lora_pattern = r'faceid.plus.sd15.lora\.safetensors$'
is_insightface = True
elif preset.startswith("faceid plus v2"):
if is_sdxl:
pattern = 'faceid.plusv2.sdxl\.(safetensors|bin)$'
lora_pattern = 'faceid.plusv2.sdxl.lora\.safetensors$'
pattern = r'faceid.plusv2.sdxl\.(safetensors|bin)$'
lora_pattern = r'faceid.plusv2.sdxl.lora\.safetensors$'
else:
pattern = 'faceid.plusv2.sd15\.(safetensors|bin)$'
lora_pattern = 'faceid.plusv2.sd15.lora\.safetensors$'
pattern = r'faceid.plusv2.sd15\.(safetensors|bin)$'
lora_pattern = r'faceid.plusv2.sd15.lora\.safetensors$'
is_insightface = True
# Community's models
elif preset.startswith("composition"):
if is_sdxl:
pattern = 'plus.composition.sdxl\.safetensors$'
pattern = r'plus.composition.sdxl\.safetensors$'
else:
pattern = 'plus.composition.sd15\.safetensors$'
pattern = r'plus.composition.sd15\.safetensors$'
else:
raise Exception(f"invalid type '{preset}'")
@@ -154,33 +154,38 @@ def insightface_loader(provider):
model.prepare(ctx_id=0, det_size=(640, 640))
return model
def encode_image_masked(clip_vision, images, mask=None):
def encode_image_masked(clip_vision, image, mask=None, batch_size=0):
model_management.load_model_gpu(clip_vision.patcher)
outputs = Output()
# Initialize lists to collect outputs
last_hidden_states = []
image_embeds = []
penultimate_hidden_states = []
if batch_size == 0:
batch_size = image.shape[0]
elif batch_size > image.shape[0]:
batch_size = image.shape[0]
# Loop over each image in the batch
for image in images:
pixel_values = clip_preprocess(image.to(clip_vision.load_device).unsqueeze(0)).float()
image_batch = torch.split(image, batch_size, dim=0)
for img in image_batch:
img = img.to(clip_vision.load_device)
pixel_values = clip_preprocess(img.to(clip_vision.load_device)).float()
# TODO: support for multiple masks
if mask is not None:
pixel_values *= mask.to(clip_vision.load_device)
pixel_values = pixel_values * mask.to(clip_vision.load_device)
out = clip_vision.model(pixel_values=pixel_values, intermediate_output=-2)
# Collect the outputs for each image
last_hidden_states.append(out[0].to(model_management.intermediate_device()))
image_embeds.append(out[2].to(model_management.intermediate_device()))
penultimate_hidden_states.append(out[1].to(model_management.intermediate_device()))
if not hasattr(outputs, "last_hidden_state"):
outputs["last_hidden_state"] = out[0].to(model_management.intermediate_device())
outputs["image_embeds"] = out[2].to(model_management.intermediate_device())
outputs["penultimate_hidden_states"] = out[1].to(model_management.intermediate_device())
else:
outputs["last_hidden_state"] = torch.cat((outputs["last_hidden_state"], out[0].to(model_management.intermediate_device())), dim=0)
outputs["image_embeds"] = torch.cat((outputs["image_embeds"], out[2].to(model_management.intermediate_device())), dim=0)
outputs["penultimate_hidden_states"] = torch.cat((outputs["penultimate_hidden_states"], out[1].to(model_management.intermediate_device())), dim=0)
# Concatenate all collected outputs across the batch
outputs = Output()
outputs["last_hidden_state"] = torch.cat(last_hidden_states, dim=0)
outputs["image_embeds"] = torch.cat(image_embeds, dim=0)
outputs["penultimate_hidden_states"] = torch.cat(penultimate_hidden_states, dim=0)
del img, pixel_values, out
return outputs
@@ -257,4 +262,4 @@ def tensor_to_image(tensor):
def image_to_tensor(image):
tensor = torch.clamp(torch.from_numpy(image).float() / 255., 0, 1)
tensor = tensor[..., [2, 1, 0]]
return tensor
return tensor