From 91a033a7605e45b7269718b6c3e416acfd39180d Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 13 Apr 2024 16:18:59 +0300 Subject: [PATCH] Add IPAdapter support --- configs/clip_vision.json | 23 ++ ip_adapter/attention_processor.py | 553 ++++++++++++++++++++++++++++++ ip_adapter/ip_adapter.py | 174 ++++++++++ ip_adapter/resampler.py | 121 +++++++ nodes.py | 147 ++++++-- 5 files changed, 982 insertions(+), 36 deletions(-) create mode 100644 configs/clip_vision.json create mode 100644 ip_adapter/attention_processor.py create mode 100644 ip_adapter/ip_adapter.py create mode 100644 ip_adapter/resampler.py diff --git a/configs/clip_vision.json b/configs/clip_vision.json new file mode 100644 index 0000000..85926d2 --- /dev/null +++ b/configs/clip_vision.json @@ -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" +} diff --git a/ip_adapter/attention_processor.py b/ip_adapter/attention_processor.py new file mode 100644 index 0000000..9b6eaed --- /dev/null +++ b/ip_adapter/attention_processor.py @@ -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 \ No newline at end of file diff --git a/ip_adapter/ip_adapter.py b/ip_adapter/ip_adapter.py new file mode 100644 index 0000000..479f43b --- /dev/null +++ b/ip_adapter/ip_adapter.py @@ -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 diff --git a/ip_adapter/resampler.py b/ip_adapter/resampler.py new file mode 100644 index 0000000..4521c8c --- /dev/null +++ b/ip_adapter/resampler.py @@ -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) \ No newline at end of file diff --git a/nodes.py b/nodes.py index 0b3ed05..de316fc 100644 --- a/nodes.py +++ b/nodes.py @@ -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)", }