Add IPAdapter support

This commit is contained in:
kijai
2024-04-13 16:18:59 +03:00
parent 2556835d0f
commit 91a033a760
5 changed files with 982 additions and 36 deletions
+23
View File
@@ -0,0 +1,23 @@
{
"_name_or_path": "",
"architectures": [
"CLIPVisionModelWithProjection"
],
"attention_dropout": 0.0,
"dropout": 0.0,
"hidden_act": "gelu",
"hidden_size": 1280,
"image_size": 224,
"initializer_factor": 1.0,
"initializer_range": 0.02,
"intermediate_size": 5120,
"layer_norm_eps": 1e-05,
"model_type": "clip_vision_model",
"num_attention_heads": 16,
"num_channels": 3,
"num_hidden_layers": 32,
"patch_size": 14,
"projection_dim": 1024,
"torch_dtype": "float16",
"transformers_version": "4.28.0.dev0"
}
+553
View File
@@ -0,0 +1,553 @@
# modified from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py
import torch
import torch.nn as nn
import torch.nn.functional as F
class AttnProcessor(nn.Module):
r"""
Default processor for performing attention-related computations.
"""
def __init__(
self,
hidden_size=None,
cross_attention_dim=None,
):
super().__init__()
def __call__(
self,
attn,
hidden_states,
encoder_hidden_states=None,
attention_mask=None,
temb=None,
):
residual = hidden_states
if attn.spatial_norm is not None:
hidden_states = attn.spatial_norm(hidden_states, temb)
input_ndim = hidden_states.ndim
if input_ndim == 4:
batch_size, channel, height, width = hidden_states.shape
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
batch_size, sequence_length, _ = (
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
)
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
if attn.group_norm is not None:
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
query = attn.to_q(hidden_states)
if encoder_hidden_states is None:
encoder_hidden_states = hidden_states
elif attn.norm_cross:
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
key = attn.to_k(encoder_hidden_states)
value = attn.to_v(encoder_hidden_states)
query = attn.head_to_batch_dim(query)
key = attn.head_to_batch_dim(key)
value = attn.head_to_batch_dim(value)
attention_probs = attn.get_attention_scores(query, key, attention_mask)
hidden_states = torch.bmm(attention_probs, value)
hidden_states = attn.batch_to_head_dim(hidden_states)
# linear proj
hidden_states = attn.to_out[0](hidden_states)
# dropout
hidden_states = attn.to_out[1](hidden_states)
if input_ndim == 4:
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
if attn.residual_connection:
hidden_states = hidden_states + residual
hidden_states = hidden_states / attn.rescale_output_factor
return hidden_states
class IPAttnProcessor(nn.Module):
r"""
Attention processor for IP-Adapater.
Args:
hidden_size (`int`):
The hidden size of the attention layer.
cross_attention_dim (`int`):
The number of channels in the `encoder_hidden_states`.
scale (`float`, defaults to 1.0):
the weight scale of image prompt.
num_tokens (`int`, defaults to 4 when do ip_adapter_plus it should be 16):
The context length of the image features.
"""
def __init__(self, hidden_size, cross_attention_dim=None, scale=1.0, num_tokens=4):
super().__init__()
self.hidden_size = hidden_size
self.cross_attention_dim = cross_attention_dim
self.scale = scale
self.num_tokens = num_tokens
self.to_k_ip = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False)
self.to_v_ip = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False)
def __call__(
self,
attn,
hidden_states,
encoder_hidden_states=None,
attention_mask=None,
temb=None,
):
residual = hidden_states
if attn.spatial_norm is not None:
hidden_states = attn.spatial_norm(hidden_states, temb)
input_ndim = hidden_states.ndim
if input_ndim == 4:
batch_size, channel, height, width = hidden_states.shape
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
batch_size, sequence_length, _ = (
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
)
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
if attn.group_norm is not None:
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
query = attn.to_q(hidden_states)
if encoder_hidden_states is None:
encoder_hidden_states = hidden_states
else:
# get encoder_hidden_states, ip_hidden_states
end_pos = encoder_hidden_states.shape[1] - self.num_tokens
encoder_hidden_states, ip_hidden_states = encoder_hidden_states[:, :end_pos, :], encoder_hidden_states[:, end_pos:, :]
if attn.norm_cross:
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
key = attn.to_k(encoder_hidden_states)
value = attn.to_v(encoder_hidden_states)
query = attn.head_to_batch_dim(query)
key = attn.head_to_batch_dim(key)
value = attn.head_to_batch_dim(value)
attention_probs = attn.get_attention_scores(query, key, attention_mask)
hidden_states = torch.bmm(attention_probs, value)
hidden_states = attn.batch_to_head_dim(hidden_states)
# for ip-adapter
ip_key = self.to_k_ip(ip_hidden_states)
ip_value = self.to_v_ip(ip_hidden_states)
ip_key = attn.head_to_batch_dim(ip_key)
ip_value = attn.head_to_batch_dim(ip_value)
ip_attention_probs = attn.get_attention_scores(query, ip_key, None)
ip_hidden_states = torch.bmm(ip_attention_probs, ip_value)
ip_hidden_states = attn.batch_to_head_dim(ip_hidden_states)
hidden_states = hidden_states + self.scale * ip_hidden_states
# linear proj
hidden_states = attn.to_out[0](hidden_states)
# dropout
hidden_states = attn.to_out[1](hidden_states)
if input_ndim == 4:
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
if attn.residual_connection:
hidden_states = hidden_states + residual
hidden_states = hidden_states / attn.rescale_output_factor
return hidden_states
class AttnProcessor2_0(torch.nn.Module):
r"""
Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0).
"""
def __init__(
self,
hidden_size=None,
cross_attention_dim=None,
):
super().__init__()
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
def __call__(
self,
attn,
hidden_states,
encoder_hidden_states=None,
attention_mask=None,
temb=None,
):
residual = hidden_states
if attn.spatial_norm is not None:
hidden_states = attn.spatial_norm(hidden_states, temb)
input_ndim = hidden_states.ndim
if input_ndim == 4:
batch_size, channel, height, width = hidden_states.shape
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
batch_size, sequence_length, _ = (
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
)
if attention_mask is not None:
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
# scaled_dot_product_attention expects attention_mask shape to be
# (batch, heads, source_length, target_length)
attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
if attn.group_norm is not None:
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
query = attn.to_q(hidden_states)
if encoder_hidden_states is None:
encoder_hidden_states = hidden_states
elif attn.norm_cross:
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
key = attn.to_k(encoder_hidden_states)
value = attn.to_v(encoder_hidden_states)
inner_dim = key.shape[-1]
head_dim = inner_dim // attn.heads
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
# the output of sdp = (batch, num_heads, seq_len, head_dim)
# TODO: add support for attn.scale when we move to Torch 2.1
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
hidden_states = hidden_states.to(query.dtype)
# linear proj
hidden_states = attn.to_out[0](hidden_states)
# dropout
hidden_states = attn.to_out[1](hidden_states)
if input_ndim == 4:
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
if attn.residual_connection:
hidden_states = hidden_states + residual
hidden_states = hidden_states / attn.rescale_output_factor
return hidden_states
class IPAttnProcessor2_0(torch.nn.Module):
r"""
Attention processor for IP-Adapater for PyTorch 2.0.
Args:
hidden_size (`int`):
The hidden size of the attention layer.
cross_attention_dim (`int`):
The number of channels in the `encoder_hidden_states`.
scale (`float`, defaults to 1.0):
the weight scale of image prompt.
num_tokens (`int`, defaults to 4 when do ip_adapter_plus it should be 16):
The context length of the image features.
"""
def __init__(self, hidden_size, cross_attention_dim=None, scale=1.0, num_tokens=4):
super().__init__()
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
self.hidden_size = hidden_size
self.cross_attention_dim = cross_attention_dim
self.scale = scale
self.num_tokens = num_tokens
self.to_k_ip = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False)
self.to_v_ip = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False)
def __call__(
self,
attn,
hidden_states,
encoder_hidden_states=None,
attention_mask=None,
temb=None,
):
residual = hidden_states
if attn.spatial_norm is not None:
hidden_states = attn.spatial_norm(hidden_states, temb)
input_ndim = hidden_states.ndim
if input_ndim == 4:
batch_size, channel, height, width = hidden_states.shape
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
batch_size, sequence_length, _ = (
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
)
if attention_mask is not None:
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
# scaled_dot_product_attention expects attention_mask shape to be
# (batch, heads, source_length, target_length)
attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
if attn.group_norm is not None:
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
query = attn.to_q(hidden_states)
if encoder_hidden_states is None:
encoder_hidden_states = hidden_states
else:
# get encoder_hidden_states, ip_hidden_states
end_pos = encoder_hidden_states.shape[1] - self.num_tokens
encoder_hidden_states, ip_hidden_states = encoder_hidden_states[:, :end_pos, :], encoder_hidden_states[:, end_pos:, :]
if attn.norm_cross:
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
key = attn.to_k(encoder_hidden_states)
value = attn.to_v(encoder_hidden_states)
inner_dim = key.shape[-1]
head_dim = inner_dim // attn.heads
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
# the output of sdp = (batch, num_heads, seq_len, head_dim)
# TODO: add support for attn.scale when we move to Torch 2.1
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
hidden_states = hidden_states.to(query.dtype)
# for ip-adapter
ip_key = self.to_k_ip(ip_hidden_states)
ip_value = self.to_v_ip(ip_hidden_states)
ip_key = ip_key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
ip_value = ip_value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
# the output of sdp = (batch, num_heads, seq_len, head_dim)
# TODO: add support for attn.scale when we move to Torch 2.1
ip_hidden_states = F.scaled_dot_product_attention(
query, ip_key, ip_value, attn_mask=None, dropout_p=0.0, is_causal=False
)
ip_hidden_states = ip_hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
ip_hidden_states = ip_hidden_states.to(query.dtype)
hidden_states = hidden_states + self.scale * ip_hidden_states
# linear proj
hidden_states = attn.to_out[0](hidden_states)
# dropout
hidden_states = attn.to_out[1](hidden_states)
if input_ndim == 4:
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
if attn.residual_connection:
hidden_states = hidden_states + residual
hidden_states = hidden_states / attn.rescale_output_factor
return hidden_states
## for controlnet
class CNAttnProcessor:
r"""
Default processor for performing attention-related computations.
"""
def __init__(self, num_tokens=4):
self.num_tokens = num_tokens
def __call__(
self,
attn,
hidden_states,
encoder_hidden_states=None,
attention_mask=None,
temb=None
):
residual = hidden_states
if attn.spatial_norm is not None:
hidden_states = attn.spatial_norm(hidden_states, temb)
input_ndim = hidden_states.ndim
if input_ndim == 4:
batch_size, channel, height, width = hidden_states.shape
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
batch_size, sequence_length, _ = (
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
)
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
if attn.group_norm is not None:
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
query = attn.to_q(hidden_states)
if encoder_hidden_states is None:
encoder_hidden_states = hidden_states
else:
end_pos = encoder_hidden_states.shape[1] - self.num_tokens
encoder_hidden_states = encoder_hidden_states[:, :end_pos] # only use text
if attn.norm_cross:
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
key = attn.to_k(encoder_hidden_states)
value = attn.to_v(encoder_hidden_states)
query = attn.head_to_batch_dim(query)
key = attn.head_to_batch_dim(key)
value = attn.head_to_batch_dim(value)
attention_probs = attn.get_attention_scores(query, key, attention_mask)
hidden_states = torch.bmm(attention_probs, value)
hidden_states = attn.batch_to_head_dim(hidden_states)
# linear proj
hidden_states = attn.to_out[0](hidden_states)
# dropout
hidden_states = attn.to_out[1](hidden_states)
if input_ndim == 4:
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
if attn.residual_connection:
hidden_states = hidden_states + residual
hidden_states = hidden_states / attn.rescale_output_factor
return hidden_states
class CNAttnProcessor2_0:
r"""
Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0).
"""
def __init__(self, num_tokens=4):
if not hasattr(F, "scaled_dot_product_attention"):
raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.")
self.num_tokens = num_tokens
def __call__(
self,
attn,
hidden_states,
encoder_hidden_states=None,
attention_mask=None,
temb=None,
):
residual = hidden_states
if attn.spatial_norm is not None:
hidden_states = attn.spatial_norm(hidden_states, temb)
input_ndim = hidden_states.ndim
if input_ndim == 4:
batch_size, channel, height, width = hidden_states.shape
hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2)
batch_size, sequence_length, _ = (
hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape
)
if attention_mask is not None:
attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)
# scaled_dot_product_attention expects attention_mask shape to be
# (batch, heads, source_length, target_length)
attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1])
if attn.group_norm is not None:
hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
query = attn.to_q(hidden_states)
if encoder_hidden_states is None:
encoder_hidden_states = hidden_states
else:
end_pos = encoder_hidden_states.shape[1] - self.num_tokens
encoder_hidden_states = encoder_hidden_states[:, :end_pos] # only use text
if attn.norm_cross:
encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states)
key = attn.to_k(encoder_hidden_states)
value = attn.to_v(encoder_hidden_states)
inner_dim = key.shape[-1]
head_dim = inner_dim // attn.heads
query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2)
# the output of sdp = (batch, num_heads, seq_len, head_dim)
# TODO: add support for attn.scale when we move to Torch 2.1
hidden_states = F.scaled_dot_product_attention(
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
)
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
hidden_states = hidden_states.to(query.dtype)
# linear proj
hidden_states = attn.to_out[0](hidden_states)
# dropout
hidden_states = attn.to_out[1](hidden_states)
if input_ndim == 4:
hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width)
if attn.residual_connection:
hidden_states = hidden_states + residual
hidden_states = hidden_states / attn.rescale_output_factor
return hidden_states
+174
View File
@@ -0,0 +1,174 @@
import torch
from diffusers.pipelines.controlnet import MultiControlNetModel
from transformers import CLIPVisionModelWithProjection, CLIPImageProcessor
from PIL import Image
if hasattr(torch.nn.functional, "scaled_dot_product_attention"):
from .attention_processor import IPAttnProcessor2_0 as IPAttnProcessor, AttnProcessor2_0 as AttnProcessor, CNAttnProcessor2_0 as CNAttnProcessor
else:
from .attention_processor import IPAttnProcessor, AttnProcessor, CNAttnProcessor
from .resampler import Resampler
class ImageProjModel(torch.nn.Module):
def __init__(self, cross_attention_dim=1024, clip_embeddings_dim=1024, clip_extra_context_tokens=4):
super().__init__()
self.cross_attention_dim = cross_attention_dim
self.clip_extra_context_tokens = clip_extra_context_tokens
self.proj = torch.nn.Linear(clip_embeddings_dim, self.clip_extra_context_tokens * cross_attention_dim)
self.norm = torch.nn.LayerNorm(cross_attention_dim)
def forward(self, image_embeds):
embeds = image_embeds
clip_extra_context_tokens = self.proj(embeds).reshape(-1, self.clip_extra_context_tokens, self.cross_attention_dim)
clip_extra_context_tokens = self.norm(clip_extra_context_tokens)
return clip_extra_context_tokens
class IPAdapter:
def __init__(self, pipe, ipadapter_ckpt_path, image_encoder_path, device="cuda", dtype=torch.float16, resample=Image.Resampling.LANCZOS):
self.pipe = pipe
self.device = device
self.dtype = dtype
# load ip adapter model
ipadapter_model = torch.load(ipadapter_ckpt_path, map_location="cpu")
# 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]
self.is_sdxl = self.output_cross_attention_dim == 2048
self.cross_attention_dim = 1280 if self.is_plus and self.is_sdxl else self.output_cross_attention_dim
self.heads = 20 if self.is_sdxl and self.is_plus else 12
self.num_tokens = 16 if self.is_plus else 4
# set image encoder
#self.image_encoder = CLIPVisionModelWithProjection.from_pretrained(image_encoder_path).to(self.device, dtype=self.dtype)
self.image_encoder = image_encoder_path
self.clip_image_processor = CLIPImageProcessor(resample=resample, do_rescale=False)
# set IPAdapter
self.set_ip_adapter()
self.image_proj_model = self.init_proj() if not self.is_plus else self.init_proj_plus()
self.image_proj_model.load_state_dict(ipadapter_model["image_proj"])
ip_layers = torch.nn.ModuleList(self.pipe.unet.attn_processors.values())
ip_layers.load_state_dict(ipadapter_model["ip_adapter"])
def init_proj(self):
image_proj_model = ImageProjModel(
cross_attention_dim=self.cross_attention_dim,
clip_embeddings_dim=self.image_encoder.config.projection_dim,
clip_extra_context_tokens=self.num_tokens,
).to(self.device, dtype=self.dtype)
return image_proj_model
def init_proj_plus(self):
image_proj_model = Resampler(
dim=self.cross_attention_dim,
depth=4,
dim_head=64,
heads=self.heads,
num_queries=self.num_tokens,
embedding_dim=self.image_encoder.config.hidden_size,
output_dim=self.output_cross_attention_dim,
ff_mult=4
).to(self.device, dtype=torch.float16)
return image_proj_model
def set_ip_adapter(self):
unet = self.pipe.unet
attn_procs = {}
for name in unet.attn_processors.keys():
cross_attention_dim = None if name.endswith("attn1.processor") else unet.config.cross_attention_dim
if name.startswith("mid_block"):
hidden_size = unet.config.block_out_channels[-1]
elif name.startswith("up_blocks"):
block_id = int(name[len("up_blocks.")])
hidden_size = list(reversed(unet.config.block_out_channels))[block_id]
elif name.startswith("down_blocks"):
block_id = int(name[len("down_blocks.")])
hidden_size = unet.config.block_out_channels[block_id]
if cross_attention_dim is None:
attn_procs[name] = AttnProcessor()
else:
attn_procs[name] = IPAttnProcessor(hidden_size=hidden_size, cross_attention_dim=cross_attention_dim).to(self.device, dtype=self.dtype)
unet.set_attn_processor(attn_procs)
if hasattr(self.pipe, "controlnet"):
if isinstance(self.pipe.controlnet, MultiControlNetModel):
for controlnet in self.pipe.controlnet.nets:
controlnet.set_attn_processor(CNAttnProcessor())
else:
self.pipe.controlnet.set_attn_processor(CNAttnProcessor())
@torch.inference_mode()
def get_image_embeds(self, images, negative_images=None):
clip_image = self.clip_image_processor(images=images, return_tensors="pt").pixel_values
clip_image = clip_image.to(self.device, dtype=torch.float16)
if not self.is_plus:
clip_image_embeds = self.image_encoder(clip_image).image_embeds
image_prompt_embeds = self.image_proj_model(clip_image_embeds)
if negative_images is not None:
negative_clip_image = self.clip_image_processor(images=negative_images, return_tensors="pt").pixel_values
negative_clip_image = negative_clip_image.to(self.device, dtype=torch.float16)
negative_image_prompt_embeds = self.image_encoder(negative_clip_image).image_embeds
else:
negative_image_prompt_embeds = torch.zeros_like(clip_image_embeds)
negative_image_prompt_embeds = self.image_proj_model(negative_image_prompt_embeds)
else:
clip_image_embeds = self.image_encoder(clip_image, output_hidden_states=True).hidden_states[-2]
image_prompt_embeds = self.image_proj_model(clip_image_embeds)
if negative_images is not None:
negative_clip_image = self.clip_image_processor(images=negative_images, return_tensors="pt").pixel_values
negative_clip_image = negative_clip_image.to(self.device, dtype=torch.float16)
negative_clip_image_embeds = self.image_encoder(negative_clip_image, output_hidden_states=True).hidden_states[-2]
else:
negative_clip_image_embeds = self.image_encoder(torch.zeros_like(clip_image), output_hidden_states=True).hidden_states[-2]
negative_image_prompt_embeds = self.image_proj_model(negative_clip_image_embeds)
num_tokens = image_prompt_embeds.shape[0] * self.num_tokens
self.set_tokens(num_tokens)
return image_prompt_embeds, negative_image_prompt_embeds
@torch.inference_mode()
def get_prompt_embeds(self, images, negative_images=None, prompt=None, negative_prompt=None, weight=[]):
prompt_embeds, negative_prompt_embeds = self.get_image_embeds(images, negative_images=negative_images)
if any(e != 1.0 for e in weight):
weight = torch.tensor(weight).unsqueeze(-1).unsqueeze(-1)
weight = weight.to(self.device)
prompt_embeds = prompt_embeds * weight
if prompt_embeds.shape[0] > 1:
prompt_embeds = torch.cat(prompt_embeds.chunk(prompt_embeds.shape[0]), dim=1)
if negative_prompt_embeds.shape[0] > 1:
negative_prompt_embeds = torch.cat(negative_prompt_embeds.chunk(negative_prompt_embeds.shape[0]), dim=1)
text_embeds = (None, None, None, None)
if prompt is not None:
text_embeds = self.pipe.encode_prompt(
prompt,
negative_prompt=negative_prompt,
device=self.device,
num_images_per_prompt=1,
do_classifier_free_guidance=True
)
prompt_embeds = torch.cat((text_embeds[0], prompt_embeds), dim=1)
negative_prompt_embeds = torch.cat((text_embeds[1], negative_prompt_embeds), dim=1)
output = (prompt_embeds, negative_prompt_embeds)
if self.is_sdxl:
output += (text_embeds[2], text_embeds[3])
return output
def set_scale(self, scale):
for attn_processor in self.pipe.unet.attn_processors.values():
if isinstance(attn_processor, IPAttnProcessor):
attn_processor.scale = scale
def set_tokens(self, num_tokens):
for attn_processor in self.pipe.unet.attn_processors.values():
if isinstance(attn_processor, IPAttnProcessor):
attn_processor.num_tokens = num_tokens
+121
View File
@@ -0,0 +1,121 @@
# modified from https://github.com/mlfoundations/open_flamingo/blob/main/open_flamingo/src/helpers.py
import math
import torch
import torch.nn as nn
# FFN
def FeedForward(dim, mult=4):
inner_dim = int(dim * mult)
return nn.Sequential(
nn.LayerNorm(dim),
nn.Linear(dim, inner_dim, bias=False),
nn.GELU(),
nn.Linear(inner_dim, dim, bias=False),
)
def reshape_tensor(x, heads):
bs, length, width = x.shape
#(bs, length, width) --> (bs, length, n_heads, dim_per_head)
x = x.view(bs, length, heads, -1)
# (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head)
x = x.transpose(1, 2)
# (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head)
x = x.reshape(bs, heads, length, -1)
return x
class PerceiverAttention(nn.Module):
def __init__(self, *, dim, dim_head=64, heads=8):
super().__init__()
self.scale = dim_head**-0.5
self.dim_head = dim_head
self.heads = heads
inner_dim = dim_head * heads
self.norm1 = nn.LayerNorm(dim)
self.norm2 = nn.LayerNorm(dim)
self.to_q = nn.Linear(dim, inner_dim, bias=False)
self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False)
self.to_out = nn.Linear(inner_dim, dim, bias=False)
def forward(self, x, latents):
"""
Args:
x (torch.Tensor): image features
shape (b, n1, D)
latent (torch.Tensor): latent features
shape (b, n2, D)
"""
x = self.norm1(x)
latents = self.norm2(latents)
b, l, _ = latents.shape
q = self.to_q(latents)
kv_input = torch.cat((x, latents), dim=-2)
k, v = self.to_kv(kv_input).chunk(2, dim=-1)
q = reshape_tensor(q, self.heads)
k = reshape_tensor(k, self.heads)
v = reshape_tensor(v, self.heads)
# attention
scale = 1 / math.sqrt(math.sqrt(self.dim_head))
weight = (q * scale) @ (k * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
out = weight @ v
out = out.permute(0, 2, 1, 3).reshape(b, l, -1)
return self.to_out(out)
class Resampler(nn.Module):
def __init__(
self,
dim=1024,
depth=8,
dim_head=64,
heads=16,
num_queries=8,
embedding_dim=768,
output_dim=1024,
ff_mult=4,
):
super().__init__()
self.latents = nn.Parameter(torch.randn(1, num_queries, dim) / dim**0.5)
self.proj_in = nn.Linear(embedding_dim, dim)
self.proj_out = nn.Linear(dim, output_dim)
self.norm_out = nn.LayerNorm(output_dim)
self.layers = nn.ModuleList([])
for _ in range(depth):
self.layers.append(
nn.ModuleList(
[
PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads),
FeedForward(dim=dim, mult=ff_mult),
]
)
)
def forward(self, x):
latents = self.latents.repeat(x.size(0), 1, 1)
x = self.proj_in(x)
for attn, ff in self.layers:
latents = attn(x, latents) + latents
latents = ff(latents) + latents
latents = self.proj_out(latents)
return self.norm_out(latents)
+111 -36
View File
@@ -25,10 +25,17 @@ try:
except:
raise ImportError("Diffusers version too old. Please update to 0.27.2 minimum.")
from .brushnet.pipeline_brushnet import StableDiffusionBrushNetPipeline
from .brushnet.brushnet import BrushNetModel
from .brushnet.unet_2d_condition import UNet2DConditionModel
from contextlib import nullcontext
from diffusers.utils import is_accelerate_available
if is_accelerate_available():
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
import safetensors.torch
from omegaconf import OmegaConf
from transformers import CLIPTokenizer
@@ -39,7 +46,7 @@ import folder_paths
script_directory = os.path.dirname(os.path.abspath(__file__))
IS_MODEL_CPU_OFFLOAD_ENABLED = False
class brushnet_model_loader:
@classmethod
def INPUT_TYPES(s):
@@ -65,6 +72,7 @@ class brushnet_model_loader:
def loadmodel(self, model, clip, vae, brushnet_model):
mm.soft_empty_cache()
dtype = mm.unet_dtype()
device = mm.get_torch_device()
custom_config = {
"model": model,
@@ -94,38 +102,44 @@ class brushnet_model_loader:
local_dir_use_symlinks=False
)
brushnet = BrushNetModel(**brushnet_config)
brushnet_sd = comfy.utils.load_torch_file(checkpoint_path)
brushnet.load_state_dict(brushnet_sd)
brushnet.to(dtype)
#create models
with (init_empty_weights() if is_accelerate_available() else nullcontext()):
brushnet = BrushNetModel(**brushnet_config)
converted_vae_config = create_vae_diffusers_config(original_config, image_size=512)
new_vae = AutoencoderKL(**converted_vae_config)
converted_unet_config = create_unet_diffusers_config(original_config, image_size=512)
new_unet = UNet2DConditionModel(**converted_unet_config)
pbar.update(1)
#load weights
brushnet_sd = comfy.utils.load_torch_file(checkpoint_path)
for key in brushnet_sd:
set_module_tensor_to_device(brushnet, key, device=device, dtype=dtype, value=brushnet_sd[key])
del brushnet_sd
clip_sd = None
load_models = [model]
load_models.append(clip.load_model())
clip_sd = clip.get_sd()
comfy.model_management.load_models_gpu(load_models)
sd = model.model.state_dict_for_saving(clip_sd, vae.get_sd(), None)
# 1. vae
converted_vae_config = create_vae_diffusers_config(original_config, image_size=512)
converted_vae = convert_ldm_vae_checkpoint(sd, converted_vae_config)
vae = AutoencoderKL(**converted_vae_config)
vae.load_state_dict(converted_vae, strict=False)
vae.to(dtype)
for key in converted_vae:
set_module_tensor_to_device(new_vae, key, device=device, dtype=dtype, value=converted_vae[key])
del converted_vae
pbar.update(1)
converted_unet = convert_ldm_unet_checkpoint(sd, converted_unet_config)
for key in converted_unet:
set_module_tensor_to_device(new_unet, key, device=device, dtype=dtype, value=converted_unet[key])
del converted_unet
pbar.update(1)
# 2. unet
converted_unet_config = create_unet_diffusers_config(original_config, image_size=512)
converted_unet = convert_ldm_unet_checkpoint(sd, converted_unet_config)
unet = UNet2DConditionModel(**converted_unet_config)
unet.load_state_dict(converted_unet, strict=False)
unet = unet.to(dtype)
pbar.update(1)
# 3. text_model
print("loading text model")
text_encoder = create_text_encoder_from_ldm_clip_checkpoint("openai/clip-vit-large-patch14",sd)
@@ -139,8 +153,8 @@ class brushnet_model_loader:
del sd
self.pipe = StableDiffusionBrushNetPipeline(
unet=unet,
vae=vae,
unet=new_unet,
vae=new_vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
scheduler=None,
@@ -255,25 +269,41 @@ class brushnet_sampler:
image = image * (1-resized_mask)
#add prompt
prompt_list = []
prompt_list.append(prompt)
if len(prompt_list) < B:
prompt_list += [prompt_list[-1]] * (B - len(prompt_list))
if 'ip_adapter' in brushnet:
print("Using IP adapter")
prompt_embeds, negative_prompt_embeds = brushnet['ip_adapter'].get_prompt_embeds(
brushnet['ip_adapter_image'],
prompt=prompt,
negative_prompt=n_prompt,
weight=[brushnet['ip_adapter_weight']]
)
use_ipadapter = True
prompt_list = None
n_prompt_list = None
else:
prompt_list = []
prompt_list.append(prompt)
if len(prompt_list) < B:
prompt_list += [prompt_list[-1]] * (B - len(prompt_list))
n_prompt_list = []
n_prompt_list.append(n_prompt)
if len(n_prompt_list) < B:
n_prompt_list += [n_prompt_list[-1]] * (B - len(n_prompt_list))
n_prompt_list = []
n_prompt_list.append(n_prompt)
if len(n_prompt_list) < B:
n_prompt_list += [n_prompt_list[-1]] * (B - len(n_prompt_list))
prompt_embeds, negative_prompt_embeds = None, None
use_ipadapter = False
#sample
generator = torch.Generator(device).manual_seed(seed)
print(prompt_list)
images = pipe(
prompt_list,
prompt=prompt_list,
negative_prompt=n_prompt_list,
image=image,
ipadapter_image=None,
prompt_embeds=prompt_embeds if use_ipadapter else None,
negative_prompt_embeds=negative_prompt_embeds if use_ipadapter else None,
mask=resized_mask,
num_inference_steps=steps,
generator=generator,
@@ -301,7 +331,7 @@ class brushnet_ella_loader:
RETURN_TYPES = ("BRUSHNET",)
RETURN_NAMES = ("brushnet",)
FUNCTION = "loadmodel"
CATEGORY = "ELLA-Wrapper"
CATEGORY = "BrushNetWrapper"
def loadmodel(self, brushnet):
print("loading ELLA")
@@ -319,7 +349,50 @@ class brushnet_ella_loader:
brushnet['pipe'].unet = ella_unet
return (brushnet,)
class brushnet_ipadapter_matteo:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"brushnet": ("BRUSHNET",),
"image": ("IMAGE",),
"ipadapter": (folder_paths.get_filename_list("ipadapter"), ),
"clip_vision" : (folder_paths.get_filename_list("clip_vision"), ),
"weight": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
},
}
RETURN_TYPES = ("BRUSHNET",)
RETURN_NAMES = ("brushnet",)
FUNCTION = "loadmodel"
CATEGORY = "BrushNetWrapper"
def loadmodel(self, image, brushnet, ipadapter, clip_vision, weight):
from .ip_adapter.ip_adapter import IPAdapter
from transformers import CLIPVisionConfig, CLIPVisionModelWithProjection
device = mm.get_torch_device()
dtype = mm.unet_dtype()
ipadapter_path = folder_paths.get_full_path("ipadapter", ipadapter)
clip_vision_path = folder_paths.get_full_path("clip_vision", clip_vision)
clip_vision_config_path = OmegaConf.load(os.path.join(script_directory, f"configs/clip_vision.json"))
clip_vision_config = CLIPVisionConfig(**clip_vision_config_path)
with (init_empty_weights() if is_accelerate_available() else nullcontext()):
image_encoder = CLIPVisionModelWithProjection(clip_vision_config)
clip_vision_sd = comfy.utils.load_torch_file(clip_vision_path)
for key in clip_vision_sd:
set_module_tensor_to_device(image_encoder, key, device=device, dtype=dtype, value=clip_vision_sd[key])
brushnet['pipe'].to(device)
ip_adapter = IPAdapter(brushnet['pipe'], ipadapter_path, image_encoder, device=device)
image = image.permute(0, 3, 1, 2).to(device)
brushnet['ip_adapter'] = ip_adapter
brushnet['ip_adapter_image'] = image
brushnet['ip_adapter_weight'] = weight
return (brushnet,)
class brushnet_sampler_ella:
@classmethod
def INPUT_TYPES(s):
@@ -447,11 +520,13 @@ NODE_CLASS_MAPPINGS = {
"brushnet_model_loader": brushnet_model_loader,
"brushnet_sampler": brushnet_sampler,
"brushnet_sampler_ella": brushnet_sampler_ella,
"brushnet_ella_loader": brushnet_ella_loader
"brushnet_ella_loader": brushnet_ella_loader,
"brushnet_ipadapter_matteo": brushnet_ipadapter_matteo,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"brushnet_model_loader": "BrushNet Model Loader",
"brushnet_sampler": "BrushNet Sampler",
"brushnet_sampler_ella": "BrushNet Sampler (ELLA)",
"brushnet_ella_loader": "BrushNet ELLA Loader"
"brushnet_ella_loader": "BrushNet ELLA Loader",
"brushnet_ipadapter_matteo": "BrushNet IP Adapter (Matteo)",
}