From aa390cbfb3d5dfe23d3ae848a499dec4449beb3b Mon Sep 17 00:00:00 2001 From: jax Date: Fri, 18 Apr 2025 13:20:24 +0800 Subject: [PATCH] add --- .DS_Store | Bin 0 -> 6148 bytes .gitignore | 3 + InstantCharacter/models/attn_processor.py | 164 ++++++ InstantCharacter/models/norm_layer.py | 46 ++ InstantCharacter/models/resampler.py | 365 ++++++++++++ InstantCharacter/models/utils.py | 139 +++++ InstantCharacter/pipeline.py | 552 ++++++++++++++++++ LICENSE | 661 ++++++++++++++++++++++ README.md | 9 + __init__.py | 19 + nodes/comfy_nodes.py | 120 ++++ requirements.txt | 10 + 12 files changed, 2088 insertions(+) create mode 100644 .DS_Store create mode 100644 .gitignore create mode 100644 InstantCharacter/models/attn_processor.py create mode 100644 InstantCharacter/models/norm_layer.py create mode 100644 InstantCharacter/models/resampler.py create mode 100644 InstantCharacter/models/utils.py create mode 100644 InstantCharacter/pipeline.py create mode 100644 LICENSE create mode 100644 README.md create mode 100644 __init__.py create mode 100644 nodes/comfy_nodes.py create mode 100644 requirements.txt diff --git a/.DS_Store b/.DS_Store new file mode 100644 index 0000000000000000000000000000000000000000..64e7a0139412be2cc36f55497ad6fe981207378e GIT binary patch literal 6148 zcmeHK&x_MQ6n@ioZE6vEP|$-A@LFmkT@<`@`{S^n9yX!}m6~jd8_Z@(k{YBGa@YUG zv;TCs{@+csDQtOEa<0=#zHYR$%!(&qYof1JqEFp(iLczZa4_r57Z#BIl5 z{q?%#+_-c1;mPc6a-PZ$N(vI#{*+xacmW>~{8FfQ zewHROeS=&cr%8_>?U<(YgnXKA%P4LmWAkS?fcTOk2%hg`^4vBvTXcdkk7z=Vfjv5> zXovB2hf${zt`+Fg@3G+1fWbQ0$BAuQJ3|LM=~|Jx*6vkF)R{woDUW8zOn_)6w%-TQKU*1GW5a5m1X nG%5-Ta~!LHkK&tfW$1G`01gdS8qor?F9J#iTUZ7Dssi5us}P`i literal 0 HcmV?d00001 diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..0637aff --- /dev/null +++ b/.gitignore @@ -0,0 +1,3 @@ +dev_notes +pushgit.bat +__pycache__ \ No newline at end of file diff --git a/InstantCharacter/models/attn_processor.py b/InstantCharacter/models/attn_processor.py new file mode 100644 index 0000000..8682bce --- /dev/null +++ b/InstantCharacter/models/attn_processor.py @@ -0,0 +1,164 @@ +from typing import Optional + +import torch.nn as nn +import torch +import torch.nn.functional as F +from diffusers.models.embeddings import apply_rotary_emb +from einops import rearrange + +from .norm_layer import RMSNorm + + +class FluxIPAttnProcessor(nn.Module): + """Attention processor used typically in processing the SD3-like self-attention projections.""" + + def __init__( + self, + hidden_size=None, + ip_hidden_states_dim=None, + ): + super().__init__() + self.norm_ip_q = RMSNorm(128, eps=1e-6) + self.to_k_ip = nn.Linear(ip_hidden_states_dim, hidden_size) + self.norm_ip_k = RMSNorm(128, eps=1e-6) + self.to_v_ip = nn.Linear(ip_hidden_states_dim, hidden_size) + + + def __call__( + self, + attn, + hidden_states: torch.FloatTensor, + encoder_hidden_states: torch.FloatTensor = None, + attention_mask: Optional[torch.FloatTensor] = None, + image_rotary_emb: Optional[torch.Tensor] = None, + emb_dict={}, + subject_emb_dict={}, + *args, + **kwargs, + ) -> torch.FloatTensor: + batch_size, _, _ = hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape + + # `sample` projections. + query = attn.to_q(hidden_states) + key = attn.to_k(hidden_states) + value = attn.to_v(hidden_states) + + # IPadapter + ip_hidden_states = self._get_ip_hidden_states( + attn, + query if encoder_hidden_states is not None else query[:, emb_dict['length_encoder_hidden_states']:], + subject_emb_dict.get('ip_hidden_states', None) + ) + + 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) + + if attn.norm_q is not None: + query = attn.norm_q(query) + if attn.norm_k is not None: + key = attn.norm_k(key) + + # the attention in FluxSingleTransformerBlock does not use `encoder_hidden_states` + if encoder_hidden_states is not None: + # `context` projections. + encoder_hidden_states_query_proj = attn.add_q_proj(encoder_hidden_states) + encoder_hidden_states_key_proj = attn.add_k_proj(encoder_hidden_states) + encoder_hidden_states_value_proj = attn.add_v_proj(encoder_hidden_states) + + encoder_hidden_states_query_proj = encoder_hidden_states_query_proj.view( + batch_size, -1, attn.heads, head_dim + ).transpose(1, 2) + encoder_hidden_states_key_proj = encoder_hidden_states_key_proj.view( + batch_size, -1, attn.heads, head_dim + ).transpose(1, 2) + encoder_hidden_states_value_proj = encoder_hidden_states_value_proj.view( + batch_size, -1, attn.heads, head_dim + ).transpose(1, 2) + + if attn.norm_added_q is not None: + encoder_hidden_states_query_proj = attn.norm_added_q(encoder_hidden_states_query_proj) + if attn.norm_added_k is not None: + encoder_hidden_states_key_proj = attn.norm_added_k(encoder_hidden_states_key_proj) + + # attention + query = torch.cat([encoder_hidden_states_query_proj, query], dim=2) + key = torch.cat([encoder_hidden_states_key_proj, key], dim=2) + value = torch.cat([encoder_hidden_states_value_proj, value], dim=2) + + if image_rotary_emb is not None: + query = apply_rotary_emb(query, image_rotary_emb) + key = apply_rotary_emb(key, image_rotary_emb) + + 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) + + + if encoder_hidden_states is not None: + encoder_hidden_states, hidden_states = ( + hidden_states[:, : encoder_hidden_states.shape[1]], + hidden_states[:, encoder_hidden_states.shape[1] :], + ) + + if ip_hidden_states is not None: + hidden_states = hidden_states + ip_hidden_states * subject_emb_dict.get('scale', 1.0) + + # linear proj + hidden_states = attn.to_out[0](hidden_states) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + encoder_hidden_states = attn.to_add_out(encoder_hidden_states) + + return hidden_states, encoder_hidden_states + else: + + if ip_hidden_states is not None: + hidden_states[:, emb_dict['length_encoder_hidden_states']:] = \ + hidden_states[:, emb_dict['length_encoder_hidden_states']:] + \ + ip_hidden_states * subject_emb_dict.get('scale', 1.0) + + return hidden_states + + + def _scaled_dot_product_attention(self, query, key, value, attention_mask=None, heads=None): + query = rearrange(query, '(b h) l c -> b h l c', h=heads) + key = rearrange(key, '(b h) l c -> b h l c', h=heads) + value = rearrange(value, '(b h) l c -> b h l c', h=heads) + hidden_states = F.scaled_dot_product_attention(query, key, value, dropout_p=0.0, is_causal=False, attn_mask=None) + hidden_states = rearrange(hidden_states, 'b h l c -> (b h) l c', h=heads) + hidden_states = hidden_states.to(query) + return hidden_states + + + def _get_ip_hidden_states( + self, + attn, + img_query, + ip_hidden_states, + ): + if ip_hidden_states is None: + return None + + if not hasattr(self, 'to_k_ip') or not hasattr(self, 'to_v_ip'): + return None + + ip_query = self.norm_ip_q(rearrange(img_query, 'b l (h d) -> b h l d', h=attn.heads)) + ip_query = rearrange(ip_query, 'b h l d -> (b h) l d') + ip_key = self.to_k_ip(ip_hidden_states) + ip_key = self.norm_ip_k(rearrange(ip_key, 'b l (h d) -> b h l d', h=attn.heads)) + ip_key = rearrange(ip_key, 'b h l d -> (b h) l d') + ip_value = self.to_v_ip(ip_hidden_states) + ip_value = attn.head_to_batch_dim(ip_value) + ip_hidden_states = self._scaled_dot_product_attention( + ip_query.to(ip_value.dtype), ip_key.to(ip_value.dtype), ip_value, None, attn.heads) + ip_hidden_states = ip_hidden_states.to(img_query.dtype) + ip_hidden_states = attn.batch_to_head_dim(ip_hidden_states) + return ip_hidden_states + diff --git a/InstantCharacter/models/norm_layer.py b/InstantCharacter/models/norm_layer.py new file mode 100644 index 0000000..ee9ff4d --- /dev/null +++ b/InstantCharacter/models/norm_layer.py @@ -0,0 +1,46 @@ +import torch.nn as nn +import torch + +class RMSNorm(nn.Module): + def __init__(self, d, p=-1., eps=1e-8, bias=False): + """ + Root Mean Square Layer Normalization + :param d: model size + :param p: partial RMSNorm, valid value [0, 1], default -1.0 (disabled) + :param eps: epsilon value, default 1e-8 + :param bias: whether use bias term for RMSNorm, disabled by + default because RMSNorm doesn't enforce re-centering invariance. + """ + super(RMSNorm, self).__init__() + + self.eps = eps + self.d = d + self.p = p + self.bias = bias + + self.scale = nn.Parameter(torch.ones(d)) + self.register_parameter("scale", self.scale) + + if self.bias: + self.offset = nn.Parameter(torch.zeros(d)) + self.register_parameter("offset", self.offset) + + def forward(self, x): + if self.p < 0. or self.p > 1.: + norm_x = x.norm(2, dim=-1, keepdim=True) + d_x = self.d + else: + partial_size = int(self.d * self.p) + partial_x, _ = torch.split(x, [partial_size, self.d - partial_size], dim=-1) + + norm_x = partial_x.norm(2, dim=-1, keepdim=True) + d_x = partial_size + + rms_x = norm_x * d_x ** (-1. / 2) + x_normed = x / (rms_x + self.eps) + + if self.bias: + return self.scale * x_normed + self.offset + + return self.scale * x_normed + diff --git a/InstantCharacter/models/resampler.py b/InstantCharacter/models/resampler.py new file mode 100644 index 0000000..a37a4ac --- /dev/null +++ b/InstantCharacter/models/resampler.py @@ -0,0 +1,365 @@ +import torch.nn as nn +import torch +import math + +from diffusers.models.transformers.transformer_2d import BasicTransformerBlock +from diffusers.models.embeddings import Timesteps, TimestepEmbedding +from timm.models.vision_transformer import Mlp + +from .norm_layer import RMSNorm + + +# 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, shift=None, scale=None): + """ + 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) + + if shift is not None and scale is not None: + latents = latents * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1) + + 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 ReshapeExpandToken(nn.Module): + def __init__(self, expand_token, token_dim): + super().__init__() + self.expand_token = expand_token + self.token_dim = token_dim + + def forward(self, x): + x = x.reshape(-1, self.expand_token, self.token_dim) + return x + + +class TimeResampler(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, + timestep_in_dim=320, + timestep_flip_sin_to_cos=True, + timestep_freq_shift=0, + expand_token=None, + extra_dim=None, + ): + super().__init__() + + self.latents = nn.Parameter(torch.randn(1, num_queries, dim) / dim**0.5) + + self.expand_token = expand_token is not None + if expand_token: + self.expand_proj = torch.nn.Sequential( + torch.nn.Linear(embedding_dim, embedding_dim * 2), + torch.nn.GELU(), + torch.nn.Linear(embedding_dim * 2, embedding_dim * expand_token), + ReshapeExpandToken(expand_token, embedding_dim), + RMSNorm(embedding_dim, eps=1e-8), + ) + + self.proj_in = nn.Linear(embedding_dim, dim) + + self.extra_feature = extra_dim is not None + if self.extra_feature: + self.proj_in_norm = RMSNorm(dim, eps=1e-8) + self.extra_proj_in = torch.nn.Sequential( + nn.Linear(extra_dim, dim), + RMSNorm(dim, eps=1e-8), + ) + + 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( + [ + # msa + PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads), + # ff + FeedForward(dim=dim, mult=ff_mult), + # adaLN + nn.Sequential(nn.SiLU(), nn.Linear(dim, 4 * dim, bias=True)) + ] + ) + ) + + # time + self.time_proj = Timesteps(timestep_in_dim, timestep_flip_sin_to_cos, timestep_freq_shift) + self.time_embedding = TimestepEmbedding(timestep_in_dim, dim, act_fn="silu") + + + def forward(self, x, timestep, need_temb=False, extra_feature=None): + timestep_emb = self.embedding_time(x, timestep) # bs, dim + + latents = self.latents.repeat(x.size(0), 1, 1) + + if self.expand_token: + x = self.expand_proj(x) + + x = self.proj_in(x) + + if self.extra_feature: + extra_feature = self.extra_proj_in(extra_feature) + x = self.proj_in_norm(x) + x = torch.cat([x, extra_feature], dim=1) + + x = x + timestep_emb[:, None] + + for attn, ff, adaLN_modulation in self.layers: + shift_msa, scale_msa, shift_mlp, scale_mlp = adaLN_modulation(timestep_emb).chunk(4, dim=1) + latents = attn(x, latents, shift_msa, scale_msa) + latents + + res = latents + for idx_ff in range(len(ff)): + layer_ff = ff[idx_ff] + latents = layer_ff(latents) + if idx_ff == 0 and isinstance(layer_ff, nn.LayerNorm): # adaLN + latents = latents * (1 + scale_mlp.unsqueeze(1)) + shift_mlp.unsqueeze(1) + latents = latents + res + + # latents = ff(latents) + latents + + latents = self.proj_out(latents) + latents = self.norm_out(latents) + + if need_temb: + return latents, timestep_emb + else: + return latents + + + def embedding_time(self, sample, timestep): + + # 1. time + timesteps = timestep + if not torch.is_tensor(timesteps): + # TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can + # This would be a good case for the `match` statement (Python 3.10+) + is_mps = sample.device.type == "mps" + if isinstance(timestep, float): + dtype = torch.float32 if is_mps else torch.float64 + else: + dtype = torch.int32 if is_mps else torch.int64 + timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device) + elif len(timesteps.shape) == 0: + timesteps = timesteps[None].to(sample.device) + + # broadcast to batch dimension in a way that's compatible with ONNX/Core ML + timesteps = timesteps.expand(sample.shape[0]) + + t_emb = self.time_proj(timesteps) + + # timesteps does not contain any weights and will always return f32 tensors + # but time_embedding might actually be running in fp16. so we need to cast here. + # there might be better ways to encapsulate this. + t_emb = t_emb.to(dtype=sample.dtype) + + emb = self.time_embedding(t_emb, None) + return emb + + +class CrossLayerCrossScaleProjector(nn.Module): + def __init__( + self, + inner_dim=2688, + num_attention_heads=42, + attention_head_dim=64, + cross_attention_dim=2688, + num_layers=4, + + # resampler + dim=1280, + depth=4, + dim_head=64, + heads=20, + num_queries=1024, + embedding_dim=1152 + 1536, + output_dim=4096, + ff_mult=4, + timestep_in_dim=320, + timestep_flip_sin_to_cos=True, + timestep_freq_shift=0, + ): + super().__init__() + + self.cross_layer_blocks = nn.ModuleList( + [ + BasicTransformerBlock( + inner_dim, + num_attention_heads, + attention_head_dim, + dropout=0, + cross_attention_dim=cross_attention_dim, + activation_fn="geglu", + num_embeds_ada_norm=None, + attention_bias=False, + only_cross_attention=False, + double_self_attention=False, + upcast_attention=False, + norm_type='layer_norm', + norm_elementwise_affine=True, + norm_eps=1e-6, + attention_type="default", + ) + for _ in range(num_layers) + ] + ) + + self.cross_scale_blocks = nn.ModuleList( + [ + BasicTransformerBlock( + inner_dim, + num_attention_heads, + attention_head_dim, + dropout=0, + cross_attention_dim=cross_attention_dim, + activation_fn="geglu", + num_embeds_ada_norm=None, + attention_bias=False, + only_cross_attention=False, + double_self_attention=False, + upcast_attention=False, + norm_type='layer_norm', + norm_elementwise_affine=True, + norm_eps=1e-6, + attention_type="default", + ) + for _ in range(num_layers) + ] + ) + + self.proj = Mlp( + in_features=inner_dim, + hidden_features=int(inner_dim*2), + act_layer=lambda: nn.GELU(approximate="tanh"), + drop=0 + ) + + self.proj_cross_layer = Mlp( + in_features=inner_dim, + hidden_features=int(inner_dim*2), + act_layer=lambda: nn.GELU(approximate="tanh"), + drop=0 + ) + + self.proj_cross_scale = Mlp( + in_features=inner_dim, + hidden_features=int(inner_dim*2), + act_layer=lambda: nn.GELU(approximate="tanh"), + drop=0 + ) + + self.resampler = TimeResampler( + dim=dim, + depth=depth, + dim_head=dim_head, + heads=heads, + num_queries=num_queries, + embedding_dim=embedding_dim, + output_dim=output_dim, + ff_mult=ff_mult, + timestep_in_dim=timestep_in_dim, + timestep_flip_sin_to_cos=timestep_flip_sin_to_cos, + timestep_freq_shift=timestep_freq_shift, + ) + + def forward(self, low_res_shallow, low_res_deep, high_res_deep, timesteps, cross_attention_kwargs=None, need_temb=True): + ''' + low_res_shallow [bs, 729*l, c] + low_res_deep [bs, 729, c] + high_res_deep [bs, 729*4, c] + ''' + + cross_layer_hidden_states = low_res_deep + for block in self.cross_layer_blocks: + cross_layer_hidden_states = block( + cross_layer_hidden_states, + encoder_hidden_states=low_res_shallow, + cross_attention_kwargs=cross_attention_kwargs, + ) + cross_layer_hidden_states = self.proj_cross_layer(cross_layer_hidden_states) + + cross_scale_hidden_states = low_res_deep + for block in self.cross_scale_blocks: + cross_scale_hidden_states = block( + cross_scale_hidden_states, + encoder_hidden_states=high_res_deep, + cross_attention_kwargs=cross_attention_kwargs, + ) + cross_scale_hidden_states = self.proj_cross_scale(cross_scale_hidden_states) + + hidden_states = self.proj(low_res_deep) + cross_scale_hidden_states + hidden_states = torch.cat([hidden_states, cross_layer_hidden_states], dim=1) + + hidden_states, timestep_emb = self.resampler(hidden_states, timesteps, need_temb=True) + return hidden_states, timestep_emb + diff --git a/InstantCharacter/models/utils.py b/InstantCharacter/models/utils.py new file mode 100644 index 0000000..36d57ec --- /dev/null +++ b/InstantCharacter/models/utils.py @@ -0,0 +1,139 @@ +from safetensors.torch import load_file +import torch +from tqdm import tqdm + +__all__ = [ + 'flux_load_lora' +] + + +def is_int(d): + try: + d = int(d) + return True + except Exception as e: + return False + + +def flux_load_lora(self, lora_file, lora_weight=1.0): + device = self.transformer.device + + # DiT 部分 + state_dict, network_alphas = self.lora_state_dict(lora_file, return_alphas=True) + state_dict = {k:v.to(device) for k,v in state_dict.items()} + + model = self.transformer + keys = list(state_dict.keys()) + keys = [k for k in keys if k.startswith('transformer.')] + + for k_lora in tqdm(keys, total=len(keys), desc=f"loading lora in transformer ..."): + v_lora = state_dict[k_lora] + + # 非 up 的都跳过 + if '.lora_A.weight' in k_lora: + continue + if '.alpha' in k_lora: + continue + + k_lora_name = k_lora.replace("transformer.", "") + k_lora_name = k_lora_name.replace(".lora_B.weight", "") + attr_name_list = k_lora_name.split('.') + + cur_attr = model + latest_attr_name = '' + for idx in range(0, len(attr_name_list)): + attr_name = attr_name_list[idx] + if is_int(attr_name): + cur_attr = cur_attr[int(attr_name)] + latest_attr_name = '' + else: + try: + if latest_attr_name != '': + cur_attr = cur_attr.__getattr__(f"{latest_attr_name}.{attr_name}") + else: + cur_attr = cur_attr.__getattr__(attr_name) + latest_attr_name = '' + except Exception as e: + if latest_attr_name != '': + latest_attr_name = f"{latest_attr_name}.{attr_name}" + else: + latest_attr_name = attr_name + + up_w = v_lora + down_w = state_dict[k_lora.replace('.lora_B.weight', '.lora_A.weight')] + + # 赋值 + einsum_a = f"ijabcdefg" + einsum_b = f"jkabcdefg" + einsum_res = f"ikabcdefg" + length_shape = len(up_w.shape) + einsum_str = f"{einsum_a[:length_shape]},{einsum_b[:length_shape]}->{einsum_res[:length_shape]}" + dtype = cur_attr.weight.data.dtype + d_w = torch.einsum(einsum_str, up_w.to(torch.float32), down_w.to(torch.float32)).to(dtype) + cur_attr.weight.data = cur_attr.weight.data + d_w * lora_weight + + + + # text encoder 部分 + raw_state_dict = load_file(lora_file) + raw_state_dict = {k:v.to(device) for k,v in raw_state_dict.items()} + + # text encoder + state_dict = {k:v for k,v in raw_state_dict.items() if 'lora_te1_' in k} + model = self.text_encoder + keys = list(state_dict.keys()) + keys = [k for k in keys if k.startswith('lora_te1_')] + + for k_lora in tqdm(keys, total=len(keys), desc=f"loading lora in text_encoder ..."): + v_lora = state_dict[k_lora] + + # 非 up 的都跳过 + if '.lora_down.weight' in k_lora: + continue + if '.alpha' in k_lora: + continue + + k_lora_name = k_lora.replace("lora_te1_", "") + k_lora_name = k_lora_name.replace(".lora_up.weight", "") + attr_name_list = k_lora_name.split('_') + + cur_attr = model + latest_attr_name = '' + for idx in range(0, len(attr_name_list)): + attr_name = attr_name_list[idx] + if is_int(attr_name): + cur_attr = cur_attr[int(attr_name)] + latest_attr_name = '' + else: + try: + if latest_attr_name != '': + cur_attr = cur_attr.__getattr__(f"{latest_attr_name}_{attr_name}") + else: + cur_attr = cur_attr.__getattr__(attr_name) + latest_attr_name = '' + except Exception as e: + if latest_attr_name != '': + latest_attr_name = f"{latest_attr_name}_{attr_name}" + else: + latest_attr_name = attr_name + + up_w = v_lora + down_w = state_dict[k_lora.replace('.lora_up.weight', '.lora_down.weight')] + + alpha = state_dict.get(k_lora.replace('.lora_up.weight', '.alpha'), None) + if alpha is None: + lora_scale = 1 + else: + rank = up_w.shape[1] + lora_scale = alpha / rank + + # 赋值 + einsum_a = f"ijabcdefg" + einsum_b = f"jkabcdefg" + einsum_res = f"ikabcdefg" + length_shape = len(up_w.shape) + einsum_str = f"{einsum_a[:length_shape]},{einsum_b[:length_shape]}->{einsum_res[:length_shape]}" + dtype = cur_attr.weight.data.dtype + d_w = torch.einsum(einsum_str, up_w.to(torch.float32), down_w.to(torch.float32)).to(dtype) + cur_attr.weight.data = cur_attr.weight.data + d_w * lora_scale * lora_weight + diff --git a/InstantCharacter/pipeline.py b/InstantCharacter/pipeline.py new file mode 100644 index 0000000..338d22d --- /dev/null +++ b/InstantCharacter/pipeline.py @@ -0,0 +1,552 @@ +# Copyright 2025 Tencent InstantX Team. All rights reserved. +# + +from PIL import Image +from einops import rearrange +import torch +from diffusers.pipelines.flux.pipeline_flux import * +from transformers import SiglipVisionModel, SiglipImageProcessor, AutoModel, AutoImageProcessor + +from models.attn_processor import FluxIPAttnProcessor +from models.resampler import CrossLayerCrossScaleProjector +from models.utils import flux_load_lora + + +# TODO +EXAMPLE_DOC_STRING = """ + Examples: + ```py + >>> import torch + >>> from diffusers import FluxPipeline + + >>> pipe = FluxPipeline.from_pretrained("black-forest-labs/FLUX.1-schnell", torch_dtype=torch.bfloat16) + >>> pipe.to("cuda") + >>> prompt = "A cat holding a sign that says hello world" + >>> # Depending on the variant being used, the pipeline call will slightly vary. + >>> # Refer to the pipeline documentation for more details. + >>> image = pipe(prompt, num_inference_steps=4, guidance_scale=0.0).images[0] + >>> image.save("flux.png") + ``` +""" + + +class InstantCharacterFluxPipeline(FluxPipeline): + + + @torch.inference_mode() + def encode_siglip_image_emb(self, siglip_image, device, dtype): + siglip_image = siglip_image.to(device, dtype=dtype) + res = self.siglip_image_encoder(siglip_image, output_hidden_states=True) + + siglip_image_embeds = res.last_hidden_state + + siglip_image_shallow_embeds = torch.cat([res.hidden_states[i] for i in [7, 13, 26]], dim=1) + + return siglip_image_embeds, siglip_image_shallow_embeds + + + @torch.inference_mode() + def encode_dinov2_image_emb(self, dinov2_image, device, dtype): + dinov2_image = dinov2_image.to(device, dtype=dtype) + res = self.dino_image_encoder_2(dinov2_image, output_hidden_states=True) + + dinov2_image_embeds = res.last_hidden_state[:, 1:] + + dinov2_image_shallow_embeds = torch.cat([res.hidden_states[i][:, 1:] for i in [9, 19, 29]], dim=1) + + return dinov2_image_embeds, dinov2_image_shallow_embeds + + + @torch.inference_mode() + def encode_image_emb(self, siglip_image, device, dtype): + object_image_pil = siglip_image + object_image_pil_low_res = [object_image_pil.resize((384, 384))] + object_image_pil_high_res = object_image_pil.resize((768, 768)) + object_image_pil_high_res = [ + object_image_pil_high_res.crop((0, 0, 384, 384)), + object_image_pil_high_res.crop((384, 0, 768, 384)), + object_image_pil_high_res.crop((0, 384, 384, 768)), + object_image_pil_high_res.crop((384, 384, 768, 768)), + ] + nb_split_image = len(object_image_pil_high_res) + + siglip_image_embeds = self.encode_siglip_image_emb( + self.siglip_image_processor(images=object_image_pil_low_res, return_tensors="pt").pixel_values, + device, + dtype + ) + dinov2_image_embeds = self.encode_dinov2_image_emb( + self.dino_image_processor_2(images=object_image_pil_low_res, return_tensors="pt").pixel_values, + device, + dtype + ) + + image_embeds_low_res_deep = torch.cat([siglip_image_embeds[0], dinov2_image_embeds[0]], dim=2) + image_embeds_low_res_shallow = torch.cat([siglip_image_embeds[1], dinov2_image_embeds[1]], dim=2) + + siglip_image_high_res = self.siglip_image_processor(images=object_image_pil_high_res, return_tensors="pt").pixel_values + siglip_image_high_res = siglip_image_high_res[None] + siglip_image_high_res = rearrange(siglip_image_high_res, 'b n c h w -> (b n) c h w') + siglip_image_high_res_embeds = self.encode_siglip_image_emb(siglip_image_high_res, device, dtype) + siglip_image_high_res_deep = rearrange(siglip_image_high_res_embeds[0], '(b n) l c -> b (n l) c', n=nb_split_image) + dinov2_image_high_res = self.dino_image_processor_2(images=object_image_pil_high_res, return_tensors="pt").pixel_values + dinov2_image_high_res = dinov2_image_high_res[None] + dinov2_image_high_res = rearrange(dinov2_image_high_res, 'b n c h w -> (b n) c h w') + dinov2_image_high_res_embeds = self.encode_dinov2_image_emb(dinov2_image_high_res, device, dtype) + dinov2_image_high_res_deep = rearrange(dinov2_image_high_res_embeds[0], '(b n) l c -> b (n l) c', n=nb_split_image) + image_embeds_high_res_deep = torch.cat([siglip_image_high_res_deep, dinov2_image_high_res_deep], dim=2) + + image_embeds_dict = dict( + image_embeds_low_res_shallow=image_embeds_low_res_shallow, + image_embeds_low_res_deep=image_embeds_low_res_deep, + image_embeds_high_res_deep=image_embeds_high_res_deep, + ) + return image_embeds_dict + + + @torch.inference_mode() + def init_ccp_and_attn_processor(self, *args, **kwargs): + subject_ip_adapter_path = kwargs['subject_ip_adapter_path'] + nb_token = kwargs['nb_token'] + state_dict = torch.load(subject_ip_adapter_path, map_location="cpu") + device, dtype = self.transformer.device, self.transformer.dtype + + print(f"=> init attn processor") + attn_procs = {} + for idx_attn, (name, v) in enumerate(self.transformer.attn_processors.items()): + attn_procs[name] = FluxIPAttnProcessor( + hidden_size=self.transformer.config.attention_head_dim * self.transformer.config.num_attention_heads, + ip_hidden_states_dim=self.text_encoder_2.config.d_model, + ).to(device, dtype=dtype) + self.transformer.set_attn_processor(attn_procs) + tmp_ip_layers = torch.nn.ModuleList(self.transformer.attn_processors.values()) + key_name = tmp_ip_layers.load_state_dict(state_dict["ip_adapter"], strict=False) + print(f"=> load attn processor: {key_name}") + + print(f"=> init project") + image_proj_model = CrossLayerCrossScaleProjector( + inner_dim=1152 + 1536, + num_attention_heads=42, + attention_head_dim=64, + cross_attention_dim=1152 + 1536, + num_layers=4, + dim=1280, + depth=4, + dim_head=64, + heads=20, + num_queries=nb_token, + embedding_dim=1152 + 1536, + output_dim=4096, + ff_mult=4, + timestep_in_dim=320, + timestep_flip_sin_to_cos=True, + timestep_freq_shift=0, + ) + image_proj_model.eval() + image_proj_model.to(device, dtype=dtype) + + key_name = image_proj_model.load_state_dict(state_dict["image_proj"], strict=False) + print(f"=> load project: {key_name}") + self.subject_image_proj_model = image_proj_model + + + @torch.inference_mode() + def init_adapter( + self, + image_encoder_path=None, + cache_dir=None, + image_encoder_2_path=None, + cache_dir_2=None, + subject_ipadapter_cfg=None, + ): + device, dtype = self.transformer.device, self.transformer.dtype + + # image encoder + print(f"=> loading image_encoder_1: {image_encoder_path}") + image_encoder = SiglipVisionModel.from_pretrained(image_encoder_path, cache_dir=cache_dir) + image_processor = SiglipImageProcessor.from_pretrained(image_encoder_path, cache_dir=cache_dir) + image_encoder.eval() + image_encoder.to(device, dtype=dtype) + self.siglip_image_encoder = image_encoder + self.siglip_image_processor = image_processor + + # image encoder 2 + print(f"=> loading image_encoder_2: {image_encoder_2_path}") + image_encoder_2 = AutoModel.from_pretrained(image_encoder_2_path, cache_dir=cache_dir_2) + image_processor_2 = AutoImageProcessor.from_pretrained(image_encoder_2_path, cache_dir=cache_dir_2) + image_encoder_2.eval() + image_encoder_2.to(device, dtype=dtype) + image_processor_2.crop_size = dict(height=384, width=384) + image_processor_2.size = dict(shortest_edge=384) + self.dino_image_encoder_2 = image_encoder_2 + self.dino_image_processor_2 = image_processor_2 + + # ccp and adapter + self.init_ccp_and_attn_processor(**subject_ipadapter_cfg) + + + @torch.no_grad() + @replace_example_docstring(EXAMPLE_DOC_STRING) + def __call__( + self, + prompt: Union[str, List[str]] = None, + prompt_2: Optional[Union[str, List[str]]] = None, + negative_prompt: Union[str, List[str]] = None, + negative_prompt_2: Optional[Union[str, List[str]]] = None, + true_cfg_scale: float = 1.0, + height: Optional[int] = None, + width: Optional[int] = None, + num_inference_steps: int = 28, + sigmas: Optional[List[float]] = None, + guidance_scale: float = 3.5, + num_images_per_prompt: Optional[int] = 1, + generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, + latents: Optional[torch.FloatTensor] = None, + prompt_embeds: Optional[torch.FloatTensor] = None, + pooled_prompt_embeds: Optional[torch.FloatTensor] = None, + ip_adapter_image: Optional[PipelineImageInput] = None, + ip_adapter_image_embeds: Optional[List[torch.Tensor]] = None, + negative_ip_adapter_image: Optional[PipelineImageInput] = None, + negative_ip_adapter_image_embeds: Optional[List[torch.Tensor]] = None, + negative_prompt_embeds: Optional[torch.FloatTensor] = None, + negative_pooled_prompt_embeds: Optional[torch.FloatTensor] = None, + output_type: Optional[str] = "pil", + return_dict: bool = True, + joint_attention_kwargs: Optional[Dict[str, Any]] = None, + callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None, + callback_on_step_end_tensor_inputs: List[str] = ["latents"], + max_sequence_length: int = 512, + subject_image: Image.Image = None, + subject_scale: float = 0.8, + + ): + r""" + Function invoked when calling the pipeline for generation. + + Args: + prompt (`str` or `List[str]`, *optional*): + The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`. + instead. + prompt_2 (`str` or `List[str]`, *optional*): + The prompt or prompts to be sent to `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is + will be used instead + height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor): + The height in pixels of the generated image. This is set to 1024 by default for the best results. + width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor): + The width in pixels of the generated image. This is set to 1024 by default for the best results. + num_inference_steps (`int`, *optional*, defaults to 50): + The number of denoising steps. More denoising steps usually lead to a higher quality image at the + expense of slower inference. + sigmas (`List[float]`, *optional*): + Custom sigmas to use for the denoising process with schedulers which support a `sigmas` argument in + their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is passed + will be used. + guidance_scale (`float`, *optional*, defaults to 7.0): + Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598). + `guidance_scale` is defined as `w` of equation 2. of [Imagen + Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale > + 1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`, + usually at the expense of lower image quality. + num_images_per_prompt (`int`, *optional*, defaults to 1): + The number of images to generate per prompt. + generator (`torch.Generator` or `List[torch.Generator]`, *optional*): + One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html) + to make generation deterministic. + latents (`torch.FloatTensor`, *optional*): + Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image + generation. Can be used to tweak the same generation with different prompts. If not provided, a latents + tensor will ge generated by sampling using the supplied random `generator`. + prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not + provided, text embeddings will be generated from `prompt` input argument. + pooled_prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. + If not provided, pooled text embeddings will be generated from `prompt` input argument. + ip_adapter_image: (`PipelineImageInput`, *optional*): Optional image input to work with IP Adapters. + ip_adapter_image_embeds (`List[torch.Tensor]`, *optional*): + Pre-generated image embeddings for IP-Adapter. It should be a list of length same as number of + IP-adapters. Each element should be a tensor of shape `(batch_size, num_images, emb_dim)`. If not + provided, embeddings are computed from the `ip_adapter_image` input argument. + negative_ip_adapter_image: + (`PipelineImageInput`, *optional*): Optional image input to work with IP Adapters. + negative_ip_adapter_image_embeds (`List[torch.Tensor]`, *optional*): + Pre-generated image embeddings for IP-Adapter. It should be a list of length same as number of + IP-adapters. Each element should be a tensor of shape `(batch_size, num_images, emb_dim)`. If not + provided, embeddings are computed from the `ip_adapter_image` input argument. + output_type (`str`, *optional*, defaults to `"pil"`): + The output format of the generate image. Choose between + [PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`. + return_dict (`bool`, *optional*, defaults to `True`): + Whether or not to return a [`~pipelines.flux.FluxPipelineOutput`] instead of a plain tuple. + joint_attention_kwargs (`dict`, *optional*): + A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under + `self.processor` in + [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). + callback_on_step_end (`Callable`, *optional*): + A function that calls at the end of each denoising steps during the inference. The function is called + with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int, + callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by + `callback_on_step_end_tensor_inputs`. + callback_on_step_end_tensor_inputs (`List`, *optional*): + The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list + will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the + `._callback_tensor_inputs` attribute of your pipeline class. + max_sequence_length (`int` defaults to 512): Maximum sequence length to use with the `prompt`. + + Examples: + + Returns: + [`~pipelines.flux.FluxPipelineOutput`] or `tuple`: [`~pipelines.flux.FluxPipelineOutput`] if `return_dict` + is True, otherwise a `tuple`. When returning a tuple, the first element is a list with the generated + images. + """ + + height = height or self.default_sample_size * self.vae_scale_factor + width = width or self.default_sample_size * self.vae_scale_factor + + # 1. Check inputs. Raise error if not correct + self.check_inputs( + prompt, + prompt_2, + height, + width, + negative_prompt=negative_prompt, + negative_prompt_2=negative_prompt_2, + prompt_embeds=prompt_embeds, + negative_prompt_embeds=negative_prompt_embeds, + pooled_prompt_embeds=pooled_prompt_embeds, + negative_pooled_prompt_embeds=negative_pooled_prompt_embeds, + callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs, + max_sequence_length=max_sequence_length, + ) + + self._guidance_scale = guidance_scale + self._joint_attention_kwargs = joint_attention_kwargs + self._interrupt = False + + # 2. Define call parameters + if prompt is not None and isinstance(prompt, str): + batch_size = 1 + elif prompt is not None and isinstance(prompt, list): + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + device = self._execution_device + dtype = self.transformer.dtype + + lora_scale = ( + self.joint_attention_kwargs.get("scale", None) if self.joint_attention_kwargs is not None else None + ) + do_true_cfg = true_cfg_scale > 1 and negative_prompt is not None + ( + prompt_embeds, + pooled_prompt_embeds, + text_ids, + ) = self.encode_prompt( + prompt=prompt, + prompt_2=prompt_2, + prompt_embeds=prompt_embeds, + pooled_prompt_embeds=pooled_prompt_embeds, + device=device, + num_images_per_prompt=num_images_per_prompt, + max_sequence_length=max_sequence_length, + lora_scale=lora_scale, + ) + if do_true_cfg: + ( + negative_prompt_embeds, + negative_pooled_prompt_embeds, + _, + ) = self.encode_prompt( + prompt=negative_prompt, + prompt_2=negative_prompt_2, + prompt_embeds=negative_prompt_embeds, + pooled_prompt_embeds=negative_pooled_prompt_embeds, + device=device, + num_images_per_prompt=num_images_per_prompt, + max_sequence_length=max_sequence_length, + lora_scale=lora_scale, + ) + + # 3.1 Prepare subject emb + if subject_image is not None: + subject_image = subject_image.resize((max(subject_image.size), max(subject_image.size))) + subject_image_embeds_dict = self.encode_image_emb(subject_image, device, dtype) + + # 4. Prepare latent variables + num_channels_latents = self.transformer.config.in_channels // 4 + latents, latent_image_ids = self.prepare_latents( + batch_size * num_images_per_prompt, + num_channels_latents, + height, + width, + prompt_embeds.dtype, + device, + generator, + latents, + ) + + # 5. Prepare timesteps + sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) if sigmas is None else sigmas + image_seq_len = latents.shape[1] + mu = calculate_shift( + image_seq_len, + self.scheduler.config.base_image_seq_len, + self.scheduler.config.max_image_seq_len, + self.scheduler.config.base_shift, + self.scheduler.config.max_shift, + ) + timesteps, num_inference_steps = retrieve_timesteps( + self.scheduler, + num_inference_steps, + device, + sigmas=sigmas, + mu=mu, + ) + num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0) + self._num_timesteps = len(timesteps) + + # handle guidance + if self.transformer.config.guidance_embeds: + guidance = torch.full([1], guidance_scale, device=device, dtype=torch.float32) + guidance = guidance.expand(latents.shape[0]) + else: + guidance = None + + if (ip_adapter_image is not None or ip_adapter_image_embeds is not None) and ( + negative_ip_adapter_image is None and negative_ip_adapter_image_embeds is None + ): + negative_ip_adapter_image = np.zeros((width, height, 3), dtype=np.uint8) + elif (ip_adapter_image is None and ip_adapter_image_embeds is None) and ( + negative_ip_adapter_image is not None or negative_ip_adapter_image_embeds is not None + ): + ip_adapter_image = np.zeros((width, height, 3), dtype=np.uint8) + + if self.joint_attention_kwargs is None: + self._joint_attention_kwargs = {} + + image_embeds = None + negative_image_embeds = None + if ip_adapter_image is not None or ip_adapter_image_embeds is not None: + image_embeds = self.prepare_ip_adapter_image_embeds( + ip_adapter_image, + ip_adapter_image_embeds, + device, + batch_size * num_images_per_prompt, + ) + if negative_ip_adapter_image is not None or negative_ip_adapter_image_embeds is not None: + negative_image_embeds = self.prepare_ip_adapter_image_embeds( + negative_ip_adapter_image, + negative_ip_adapter_image_embeds, + device, + batch_size * num_images_per_prompt, + ) + + # 6. Denoising loop + with self.progress_bar(total=num_inference_steps) as progress_bar: + for i, t in enumerate(timesteps): + if self.interrupt: + continue + + if image_embeds is not None: + self._joint_attention_kwargs["ip_adapter_image_embeds"] = image_embeds + # broadcast to batch dimension in a way that's compatible with ONNX/Core ML + timestep = t.expand(latents.shape[0]).to(latents.dtype) + + + # subject adapter + if subject_image is not None: + subject_image_prompt_embeds = self.subject_image_proj_model( + low_res_shallow=subject_image_embeds_dict['image_embeds_low_res_shallow'], + low_res_deep=subject_image_embeds_dict['image_embeds_low_res_deep'], + high_res_deep=subject_image_embeds_dict['image_embeds_high_res_deep'], + timesteps=timestep.to(dtype=latents.dtype), + need_temb=True + )[0] + self._joint_attention_kwargs['emb_dict'] = dict( + length_encoder_hidden_states=prompt_embeds.shape[1] + ) + self._joint_attention_kwargs['subject_emb_dict'] = dict( + ip_hidden_states=subject_image_prompt_embeds, + scale=subject_scale, + ) + + noise_pred = self.transformer( + hidden_states=latents, + timestep=timestep / 1000, + guidance=guidance, + pooled_projections=pooled_prompt_embeds, + encoder_hidden_states=prompt_embeds, + txt_ids=text_ids, + img_ids=latent_image_ids, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + )[0] + + if do_true_cfg: + if negative_image_embeds is not None: + self._joint_attention_kwargs["ip_adapter_image_embeds"] = negative_image_embeds + neg_noise_pred = self.transformer( + hidden_states=latents, + timestep=timestep / 1000, + guidance=guidance, + pooled_projections=negative_pooled_prompt_embeds, + encoder_hidden_states=negative_prompt_embeds, + txt_ids=text_ids, + img_ids=latent_image_ids, + joint_attention_kwargs=self.joint_attention_kwargs, + return_dict=False, + )[0] + noise_pred = neg_noise_pred + true_cfg_scale * (noise_pred - neg_noise_pred) + + # compute the previous noisy sample x_t -> x_t-1 + latents_dtype = latents.dtype + latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0] + + if latents.dtype != latents_dtype: + if torch.backends.mps.is_available(): + # some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272 + latents = latents.to(latents_dtype) + + if callback_on_step_end is not None: + callback_kwargs = {} + for k in callback_on_step_end_tensor_inputs: + callback_kwargs[k] = locals()[k] + callback_outputs = callback_on_step_end(self, i, t, callback_kwargs) + + latents = callback_outputs.pop("latents", latents) + prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds) + + # call the callback, if provided + if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0): + progress_bar.update() + + if XLA_AVAILABLE: + xm.mark_step() + + if output_type == "latent": + image = latents + + else: + latents = self._unpack_latents(latents, height, width, self.vae_scale_factor) + latents = (latents / self.vae.config.scaling_factor) + self.vae.config.shift_factor + image = self.vae.decode(latents, return_dict=False)[0] + image = self.image_processor.postprocess(image, output_type=output_type) + + # Offload all models + self.maybe_free_model_hooks() + + if not return_dict: + return (image,) + + return FluxPipelineOutput(images=image) + + + def with_style_lora(self, lora_file_path, lora_weight=1.0, trigger='', *args, **kwargs): + flux_load_lora(self, lora_file_path, lora_weight) + kwargs['prompt'] = f"{trigger}, {kwargs['prompt']}" + res = self.__call__(*args, **kwargs) + flux_load_lora(self, lora_file_path, -lora_weight) + return res + diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..0ad25db --- /dev/null +++ b/LICENSE @@ -0,0 +1,661 @@ + GNU AFFERO GENERAL PUBLIC LICENSE + Version 3, 19 November 2007 + + Copyright (C) 2007 Free Software Foundation, Inc. + Everyone is permitted to copy and distribute verbatim copies + of this license document, but changing it is not allowed. + + Preamble + + The GNU Affero General Public License is a free, copyleft license for +software and other kinds of works, specifically designed to ensure +cooperation with the community in the case of network server software. + + The licenses for most software and other practical works are designed +to take away your freedom to share and change the works. By contrast, +our General Public Licenses are intended to guarantee your freedom to +share and change all versions of a program--to make sure it remains free +software for all its users. + + When we speak of free software, we are referring to freedom, not +price. Our General Public Licenses are designed to make sure that you +have the freedom to distribute copies of free software (and charge for +them if you wish), that you receive source code or can get it if you +want it, that you can change the software or use pieces of it in new +free programs, and that you know you can do these things. + + Developers that use our General Public Licenses protect your rights +with two steps: (1) assert copyright on the software, and (2) offer +you this License which gives you legal permission to copy, distribute +and/or modify the software. + + A secondary benefit of defending all users' freedom is that +improvements made in alternate versions of the program, if they +receive widespread use, become available for other developers to +incorporate. Many developers of free software are heartened and +encouraged by the resulting cooperation. However, in the case of +software used on network servers, this result may fail to come about. +The GNU General Public License permits making a modified version and +letting the public access it on a server without ever releasing its +source code to the public. + + The GNU Affero General Public License is designed specifically to +ensure that, in such cases, the modified source code becomes available +to the community. It requires the operator of a network server to +provide the source code of the modified version running there to the +users of that server. Therefore, public use of a modified version, on +a publicly accessible server, gives the public access to the source +code of the modified version. + + An older license, called the Affero General Public License and +published by Affero, was designed to accomplish similar goals. This is +a different license, not a version of the Affero GPL, but Affero has +released a new version of the Affero GPL which permits relicensing under +this license. + + The precise terms and conditions for copying, distribution and +modification follow. + + TERMS AND CONDITIONS + + 0. Definitions. + + "This License" refers to version 3 of the GNU Affero General Public License. + + "Copyright" also means copyright-like laws that apply to other kinds of +works, such as semiconductor masks. + + "The Program" refers to any copyrightable work licensed under this +License. Each licensee is addressed as "you". "Licensees" and +"recipients" may be individuals or organizations. + + To "modify" a work means to copy from or adapt all or part of the work +in a fashion requiring copyright permission, other than the making of an +exact copy. The resulting work is called a "modified version" of the +earlier work or a work "based on" the earlier work. + + A "covered work" means either the unmodified Program or a work based +on the Program. + + To "propagate" a work means to do anything with it that, without +permission, would make you directly or secondarily liable for +infringement under applicable copyright law, except executing it on a +computer or modifying a private copy. Propagation includes copying, +distribution (with or without modification), making available to the +public, and in some countries other activities as well. + + To "convey" a work means any kind of propagation that enables other +parties to make or receive copies. Mere interaction with a user through +a computer network, with no transfer of a copy, is not conveying. + + An interactive user interface displays "Appropriate Legal Notices" +to the extent that it includes a convenient and prominently visible +feature that (1) displays an appropriate copyright notice, and (2) +tells the user that there is no warranty for the work (except to the +extent that warranties are provided), that licensees may convey the +work under this License, and how to view a copy of this License. If +the interface presents a list of user commands or options, such as a +menu, a prominent item in the list meets this criterion. + + 1. Source Code. + + The "source code" for a work means the preferred form of the work +for making modifications to it. "Object code" means any non-source +form of a work. + + A "Standard Interface" means an interface that either is an official +standard defined by a recognized standards body, or, in the case of +interfaces specified for a particular programming language, one that +is widely used among developers working in that language. + + The "System Libraries" of an executable work include anything, other +than the work as a whole, that (a) is included in the normal form of +packaging a Major Component, but which is not part of that Major +Component, and (b) serves only to enable use of the work with that +Major Component, or to implement a Standard Interface for which an +implementation is available to the public in source code form. A +"Major Component", in this context, means a major essential component +(kernel, window system, and so on) of the specific operating system +(if any) on which the executable work runs, or a compiler used to +produce the work, or an object code interpreter used to run it. + + The "Corresponding Source" for a work in object code form means all +the source code needed to generate, install, and (for an executable +work) run the object code and to modify the work, including scripts to +control those activities. However, it does not include the work's +System Libraries, or general-purpose tools or generally available free +programs which are used unmodified in performing those activities but +which are not part of the work. For example, Corresponding Source +includes interface definition files associated with source files for +the work, and the source code for shared libraries and dynamically +linked subprograms that the work is specifically designed to require, +such as by intimate data communication or control flow between those +subprograms and other parts of the work. + + The Corresponding Source need not include anything that users +can regenerate automatically from other parts of the Corresponding +Source. + + The Corresponding Source for a work in source code form is that +same work. + + 2. Basic Permissions. + + All rights granted under this License are granted for the term of +copyright on the Program, and are irrevocable provided the stated +conditions are met. This License explicitly affirms your unlimited +permission to run the unmodified Program. The output from running a +covered work is covered by this License only if the output, given its +content, constitutes a covered work. This License acknowledges your +rights of fair use or other equivalent, as provided by copyright law. + + You may make, run and propagate covered works that you do not +convey, without conditions so long as your license otherwise remains +in force. You may convey covered works to others for the sole purpose +of having them make modifications exclusively for you, or provide you +with facilities for running those works, provided that you comply with +the terms of this License in conveying all material for which you do +not control copyright. Those thus making or running the covered works +for you must do so exclusively on your behalf, under your direction +and control, on terms that prohibit them from making any copies of +your copyrighted material outside their relationship with you. + + Conveying under any other circumstances is permitted solely under +the conditions stated below. Sublicensing is not allowed; section 10 +makes it unnecessary. + + 3. Protecting Users' Legal Rights From Anti-Circumvention Law. + + No covered work shall be deemed part of an effective technological +measure under any applicable law fulfilling obligations under article +11 of the WIPO copyright treaty adopted on 20 December 1996, or +similar laws prohibiting or restricting circumvention of such +measures. + + When you convey a covered work, you waive any legal power to forbid +circumvention of technological measures to the extent such circumvention +is effected by exercising rights under this License with respect to +the covered work, and you disclaim any intention to limit operation or +modification of the work as a means of enforcing, against the work's +users, your or third parties' legal rights to forbid circumvention of +technological measures. + + 4. Conveying Verbatim Copies. + + You may convey verbatim copies of the Program's source code as you +receive it, in any medium, provided that you conspicuously and +appropriately publish on each copy an appropriate copyright notice; +keep intact all notices stating that this License and any +non-permissive terms added in accord with section 7 apply to the code; +keep intact all notices of the absence of any warranty; and give all +recipients a copy of this License along with the Program. + + You may charge any price or no price for each copy that you convey, +and you may offer support or warranty protection for a fee. + + 5. Conveying Modified Source Versions. + + You may convey a work based on the Program, or the modifications to +produce it from the Program, in the form of source code under the +terms of section 4, provided that you also meet all of these conditions: + + a) The work must carry prominent notices stating that you modified + it, and giving a relevant date. + + b) The work must carry prominent notices stating that it is + released under this License and any conditions added under section + 7. This requirement modifies the requirement in section 4 to + "keep intact all notices". + + c) You must license the entire work, as a whole, under this + License to anyone who comes into possession of a copy. This + License will therefore apply, along with any applicable section 7 + additional terms, to the whole of the work, and all its parts, + regardless of how they are packaged. This License gives no + permission to license the work in any other way, but it does not + invalidate such permission if you have separately received it. + + d) If the work has interactive user interfaces, each must display + Appropriate Legal Notices; however, if the Program has interactive + interfaces that do not display Appropriate Legal Notices, your + work need not make them do so. + + A compilation of a covered work with other separate and independent +works, which are not by their nature extensions of the covered work, +and which are not combined with it such as to form a larger program, +in or on a volume of a storage or distribution medium, is called an +"aggregate" if the compilation and its resulting copyright are not +used to limit the access or legal rights of the compilation's users +beyond what the individual works permit. Inclusion of a covered work +in an aggregate does not cause this License to apply to the other +parts of the aggregate. + + 6. Conveying Non-Source Forms. + + You may convey a covered work in object code form under the terms +of sections 4 and 5, provided that you also convey the +machine-readable Corresponding Source under the terms of this License, +in one of these ways: + + a) Convey the object code in, or embodied in, a physical product + (including a physical distribution medium), accompanied by the + Corresponding Source fixed on a durable physical medium + customarily used for software interchange. + + b) Convey the object code in, or embodied in, a physical product + (including a physical distribution medium), accompanied by a + written offer, valid for at least three years and valid for as + long as you offer spare parts or customer support for that product + model, to give anyone who possesses the object code either (1) a + copy of the Corresponding Source for all the software in the + product that is covered by this License, on a durable physical + medium customarily used for software interchange, for a price no + more than your reasonable cost of physically performing this + conveying of source, or (2) access to copy the + Corresponding Source from a network server at no charge. + + c) Convey individual copies of the object code with a copy of the + written offer to provide the Corresponding Source. This + alternative is allowed only occasionally and noncommercially, and + only if you received the object code with such an offer, in accord + with subsection 6b. + + d) Convey the object code by offering access from a designated + place (gratis or for a charge), and offer equivalent access to the + Corresponding Source in the same way through the same place at no + further charge. You need not require recipients to copy the + Corresponding Source along with the object code. If the place to + copy the object code is a network server, the Corresponding Source + may be on a different server (operated by you or a third party) + that supports equivalent copying facilities, provided you maintain + clear directions next to the object code saying where to find the + Corresponding Source. Regardless of what server hosts the + Corresponding Source, you remain obligated to ensure that it is + available for as long as needed to satisfy these requirements. + + e) Convey the object code using peer-to-peer transmission, provided + you inform other peers where the object code and Corresponding + Source of the work are being offered to the general public at no + charge under subsection 6d. + + A separable portion of the object code, whose source code is excluded +from the Corresponding Source as a System Library, need not be +included in conveying the object code work. + + A "User Product" is either (1) a "consumer product", which means any +tangible personal property which is normally used for personal, family, +or household purposes, or (2) anything designed or sold for incorporation +into a dwelling. In determining whether a product is a consumer product, +doubtful cases shall be resolved in favor of coverage. For a particular +product received by a particular user, "normally used" refers to a +typical or common use of that class of product, regardless of the status +of the particular user or of the way in which the particular user +actually uses, or expects or is expected to use, the product. A product +is a consumer product regardless of whether the product has substantial +commercial, industrial or non-consumer uses, unless such uses represent +the only significant mode of use of the product. + + "Installation Information" for a User Product means any methods, +procedures, authorization keys, or other information required to install +and execute modified versions of a covered work in that User Product from +a modified version of its Corresponding Source. The information must +suffice to ensure that the continued functioning of the modified object +code is in no case prevented or interfered with solely because +modification has been made. + + If you convey an object code work under this section in, or with, or +specifically for use in, a User Product, and the conveying occurs as +part of a transaction in which the right of possession and use of the +User Product is transferred to the recipient in perpetuity or for a +fixed term (regardless of how the transaction is characterized), the +Corresponding Source conveyed under this section must be accompanied +by the Installation Information. But this requirement does not apply +if neither you nor any third party retains the ability to install +modified object code on the User Product (for example, the work has +been installed in ROM). + + The requirement to provide Installation Information does not include a +requirement to continue to provide support service, warranty, or updates +for a work that has been modified or installed by the recipient, or for +the User Product in which it has been modified or installed. Access to a +network may be denied when the modification itself materially and +adversely affects the operation of the network or violates the rules and +protocols for communication across the network. + + Corresponding Source conveyed, and Installation Information provided, +in accord with this section must be in a format that is publicly +documented (and with an implementation available to the public in +source code form), and must require no special password or key for +unpacking, reading or copying. + + 7. Additional Terms. + + "Additional permissions" are terms that supplement the terms of this +License by making exceptions from one or more of its conditions. +Additional permissions that are applicable to the entire Program shall +be treated as though they were included in this License, to the extent +that they are valid under applicable law. If additional permissions +apply only to part of the Program, that part may be used separately +under those permissions, but the entire Program remains governed by +this License without regard to the additional permissions. + + When you convey a copy of a covered work, you may at your option +remove any additional permissions from that copy, or from any part of +it. (Additional permissions may be written to require their own +removal in certain cases when you modify the work.) You may place +additional permissions on material, added by you to a covered work, +for which you have or can give appropriate copyright permission. + + Notwithstanding any other provision of this License, for material you +add to a covered work, you may (if authorized by the copyright holders of +that material) supplement the terms of this License with terms: + + a) Disclaiming warranty or limiting liability differently from the + terms of sections 15 and 16 of this License; or + + b) Requiring preservation of specified reasonable legal notices or + author attributions in that material or in the Appropriate Legal + Notices displayed by works containing it; or + + c) Prohibiting misrepresentation of the origin of that material, or + requiring that modified versions of such material be marked in + reasonable ways as different from the original version; or + + d) Limiting the use for publicity purposes of names of licensors or + authors of the material; or + + e) Declining to grant rights under trademark law for use of some + trade names, trademarks, or service marks; or + + f) Requiring indemnification of licensors and authors of that + material by anyone who conveys the material (or modified versions of + it) with contractual assumptions of liability to the recipient, for + any liability that these contractual assumptions directly impose on + those licensors and authors. + + All other non-permissive additional terms are considered "further +restrictions" within the meaning of section 10. If the Program as you +received it, or any part of it, contains a notice stating that it is +governed by this License along with a term that is a further +restriction, you may remove that term. If a license document contains +a further restriction but permits relicensing or conveying under this +License, you may add to a covered work material governed by the terms +of that license document, provided that the further restriction does +not survive such relicensing or conveying. + + If you add terms to a covered work in accord with this section, you +must place, in the relevant source files, a statement of the +additional terms that apply to those files, or a notice indicating +where to find the applicable terms. + + Additional terms, permissive or non-permissive, may be stated in the +form of a separately written license, or stated as exceptions; +the above requirements apply either way. + + 8. Termination. + + You may not propagate or modify a covered work except as expressly +provided under this License. Any attempt otherwise to propagate or +modify it is void, and will automatically terminate your rights under +this License (including any patent licenses granted under the third +paragraph of section 11). + + However, if you cease all violation of this License, then your +license from a particular copyright holder is reinstated (a) +provisionally, unless and until the copyright holder explicitly and +finally terminates your license, and (b) permanently, if the copyright +holder fails to notify you of the violation by some reasonable means +prior to 60 days after the cessation. + + Moreover, your license from a particular copyright holder is +reinstated permanently if the copyright holder notifies you of the +violation by some reasonable means, this is the first time you have +received notice of violation of this License (for any work) from that +copyright holder, and you cure the violation prior to 30 days after +your receipt of the notice. + + Termination of your rights under this section does not terminate the +licenses of parties who have received copies or rights from you under +this License. If your rights have been terminated and not permanently +reinstated, you do not qualify to receive new licenses for the same +material under section 10. + + 9. Acceptance Not Required for Having Copies. + + You are not required to accept this License in order to receive or +run a copy of the Program. Ancillary propagation of a covered work +occurring solely as a consequence of using peer-to-peer transmission +to receive a copy likewise does not require acceptance. However, +nothing other than this License grants you permission to propagate or +modify any covered work. These actions infringe copyright if you do +not accept this License. Therefore, by modifying or propagating a +covered work, you indicate your acceptance of this License to do so. + + 10. Automatic Licensing of Downstream Recipients. + + Each time you convey a covered work, the recipient automatically +receives a license from the original licensors, to run, modify and +propagate that work, subject to this License. You are not responsible +for enforcing compliance by third parties with this License. + + An "entity transaction" is a transaction transferring control of an +organization, or substantially all assets of one, or subdividing an +organization, or merging organizations. If propagation of a covered +work results from an entity transaction, each party to that +transaction who receives a copy of the work also receives whatever +licenses to the work the party's predecessor in interest had or could +give under the previous paragraph, plus a right to possession of the +Corresponding Source of the work from the predecessor in interest, if +the predecessor has it or can get it with reasonable efforts. + + You may not impose any further restrictions on the exercise of the +rights granted or affirmed under this License. For example, you may +not impose a license fee, royalty, or other charge for exercise of +rights granted under this License, and you may not initiate litigation +(including a cross-claim or counterclaim in a lawsuit) alleging that +any patent claim is infringed by making, using, selling, offering for +sale, or importing the Program or any portion of it. + + 11. Patents. + + A "contributor" is a copyright holder who authorizes use under this +License of the Program or a work on which the Program is based. The +work thus licensed is called the contributor's "contributor version". + + A contributor's "essential patent claims" are all patent claims +owned or controlled by the contributor, whether already acquired or +hereafter acquired, that would be infringed by some manner, permitted +by this License, of making, using, or selling its contributor version, +but do not include claims that would be infringed only as a +consequence of further modification of the contributor version. For +purposes of this definition, "control" includes the right to grant +patent sublicenses in a manner consistent with the requirements of +this License. + + Each contributor grants you a non-exclusive, worldwide, royalty-free +patent license under the contributor's essential patent claims, to +make, use, sell, offer for sale, import and otherwise run, modify and +propagate the contents of its contributor version. + + In the following three paragraphs, a "patent license" is any express +agreement or commitment, however denominated, not to enforce a patent +(such as an express permission to practice a patent or covenant not to +sue for patent infringement). To "grant" such a patent license to a +party means to make such an agreement or commitment not to enforce a +patent against the party. + + If you convey a covered work, knowingly relying on a patent license, +and the Corresponding Source of the work is not available for anyone +to copy, free of charge and under the terms of this License, through a +publicly available network server or other readily accessible means, +then you must either (1) cause the Corresponding Source to be so +available, or (2) arrange to deprive yourself of the benefit of the +patent license for this particular work, or (3) arrange, in a manner +consistent with the requirements of this License, to extend the patent +license to downstream recipients. "Knowingly relying" means you have +actual knowledge that, but for the patent license, your conveying the +covered work in a country, or your recipient's use of the covered work +in a country, would infringe one or more identifiable patents in that +country that you have reason to believe are valid. + + If, pursuant to or in connection with a single transaction or +arrangement, you convey, or propagate by procuring conveyance of, a +covered work, and grant a patent license to some of the parties +receiving the covered work authorizing them to use, propagate, modify +or convey a specific copy of the covered work, then the patent license +you grant is automatically extended to all recipients of the covered +work and works based on it. + + A patent license is "discriminatory" if it does not include within +the scope of its coverage, prohibits the exercise of, or is +conditioned on the non-exercise of one or more of the rights that are +specifically granted under this License. You may not convey a covered +work if you are a party to an arrangement with a third party that is +in the business of distributing software, under which you make payment +to the third party based on the extent of your activity of conveying +the work, and under which the third party grants, to any of the +parties who would receive the covered work from you, a discriminatory +patent license (a) in connection with copies of the covered work +conveyed by you (or copies made from those copies), or (b) primarily +for and in connection with specific products or compilations that +contain the covered work, unless you entered into that arrangement, +or that patent license was granted, prior to 28 March 2007. + + Nothing in this License shall be construed as excluding or limiting +any implied license or other defenses to infringement that may +otherwise be available to you under applicable patent law. + + 12. No Surrender of Others' Freedom. + + If conditions are imposed on you (whether by court order, agreement or +otherwise) that contradict the conditions of this License, they do not +excuse you from the conditions of this License. If you cannot convey a +covered work so as to satisfy simultaneously your obligations under this +License and any other pertinent obligations, then as a consequence you may +not convey it at all. For example, if you agree to terms that obligate you +to collect a royalty for further conveying from those to whom you convey +the Program, the only way you could satisfy both those terms and this +License would be to refrain entirely from conveying the Program. + + 13. Remote Network Interaction; Use with the GNU General Public License. + + Notwithstanding any other provision of this License, if you modify the +Program, your modified version must prominently offer all users +interacting with it remotely through a computer network (if your version +supports such interaction) an opportunity to receive the Corresponding +Source of your version by providing access to the Corresponding Source +from a network server at no charge, through some standard or customary +means of facilitating copying of software. This Corresponding Source +shall include the Corresponding Source for any work covered by version 3 +of the GNU General Public License that is incorporated pursuant to the +following paragraph. + + Notwithstanding any other provision of this License, you have +permission to link or combine any covered work with a work licensed +under version 3 of the GNU General Public License into a single +combined work, and to convey the resulting work. The terms of this +License will continue to apply to the part which is the covered work, +but the work with which it is combined will remain governed by version +3 of the GNU General Public License. + + 14. Revised Versions of this License. + + The Free Software Foundation may publish revised and/or new versions of +the GNU Affero General Public License from time to time. Such new versions +will be similar in spirit to the present version, but may differ in detail to +address new problems or concerns. + + Each version is given a distinguishing version number. If the +Program specifies that a certain numbered version of the GNU Affero General +Public License "or any later version" applies to it, you have the +option of following the terms and conditions either of that numbered +version or of any later version published by the Free Software +Foundation. If the Program does not specify a version number of the +GNU Affero General Public License, you may choose any version ever published +by the Free Software Foundation. + + If the Program specifies that a proxy can decide which future +versions of the GNU Affero General Public License can be used, that proxy's +public statement of acceptance of a version permanently authorizes you +to choose that version for the Program. + + Later license versions may give you additional or different +permissions. However, no additional obligations are imposed on any +author or copyright holder as a result of your choosing to follow a +later version. + + 15. Disclaimer of Warranty. + + THERE IS NO WARRANTY FOR THE PROGRAM, TO THE EXTENT PERMITTED BY +APPLICABLE LAW. EXCEPT WHEN OTHERWISE STATED IN WRITING THE COPYRIGHT +HOLDERS AND/OR OTHER PARTIES PROVIDE THE PROGRAM "AS IS" WITHOUT WARRANTY +OF ANY KIND, EITHER EXPRESSED OR IMPLIED, INCLUDING, BUT NOT LIMITED TO, +THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR +PURPOSE. THE ENTIRE RISK AS TO THE QUALITY AND PERFORMANCE OF THE PROGRAM +IS WITH YOU. SHOULD THE PROGRAM PROVE DEFECTIVE, YOU ASSUME THE COST OF +ALL NECESSARY SERVICING, REPAIR OR CORRECTION. + + 16. Limitation of Liability. + + IN NO EVENT UNLESS REQUIRED BY APPLICABLE LAW OR AGREED TO IN WRITING +WILL ANY COPYRIGHT HOLDER, OR ANY OTHER PARTY WHO MODIFIES AND/OR CONVEYS +THE PROGRAM AS PERMITTED ABOVE, BE LIABLE TO YOU FOR DAMAGES, INCLUDING ANY +GENERAL, SPECIAL, INCIDENTAL OR CONSEQUENTIAL DAMAGES ARISING OUT OF THE +USE OR INABILITY TO USE THE PROGRAM (INCLUDING BUT NOT LIMITED TO LOSS OF +DATA OR DATA BEING RENDERED INACCURATE OR LOSSES SUSTAINED BY YOU OR THIRD +PARTIES OR A FAILURE OF THE PROGRAM TO OPERATE WITH ANY OTHER PROGRAMS), +EVEN IF SUCH HOLDER OR OTHER PARTY HAS BEEN ADVISED OF THE POSSIBILITY OF +SUCH DAMAGES. + + 17. Interpretation of Sections 15 and 16. + + If the disclaimer of warranty and limitation of liability provided +above cannot be given local legal effect according to their terms, +reviewing courts shall apply local law that most closely approximates +an absolute waiver of all civil liability in connection with the +Program, unless a warranty or assumption of liability accompanies a +copy of the Program in return for a fee. + + END OF TERMS AND CONDITIONS + + How to Apply These Terms to Your New Programs + + If you develop a new program, and you want it to be of the greatest +possible use to the public, the best way to achieve this is to make it +free software which everyone can redistribute and change under these terms. + + To do so, attach the following notices to the program. It is safest +to attach them to the start of each source file to most effectively +state the exclusion of warranty; and each file should have at least +the "copyright" line and a pointer to where the full notice is found. + + + Copyright (C) + + This program is free software: you can redistribute it and/or modify + it under the terms of the GNU Affero General Public License as published + by the Free Software Foundation, either version 3 of the License, or + (at your option) any later version. + + This program is distributed in the hope that it will be useful, + but WITHOUT ANY WARRANTY; without even the implied warranty of + MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + GNU Affero General Public License for more details. + + You should have received a copy of the GNU Affero General Public License + along with this program. If not, see . + +Also add information on how to contact you by electronic and paper mail. + + If your software can interact with users remotely through a computer +network, you should also make sure that it provides a way for users to +get its source. For example, if your program is a web application, its +interface could display a "Source" link that leads users to an archive +of the code. There are many ways you could offer source, and different +solutions will be better for different programs; see section 13 for the +specific requirements. + + You should also get your employer (if you work as a programmer) or school, +if any, to sign a "copyright disclaimer" for the program, if necessary. +For more information on this, and how to apply and follow the GNU AGPL, see +. diff --git a/README.md b/README.md new file mode 100644 index 0000000..ac2cb05 --- /dev/null +++ b/README.md @@ -0,0 +1,9 @@ +# comfyui-model-dynamic-loader + +for comfyonline dynamic loader + +https://www.comfyonline.app +comfyonline is comfyui cloud website, Run ComfyUI workflows online and deploy APIs with one click + +Provides an online environment for running your ComfyUI workflows, with the ability to generate APIs for easy AI application development. + diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..9db9a81 --- /dev/null +++ b/__init__.py @@ -0,0 +1,19 @@ + + + +# 注册节点 +from .nodes.comfy_nodes import InstantCharacterLoadModel, InstantCharacterGenerate + + +NODE_CLASS_MAPPINGS = { + "InstantCharacterLoadModel": InstantCharacterLoadModel, + "InstantCharacterGenerate": InstantCharacterGenerate, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "InstantCharacterLoadModel": "InstantCharacter Load Model", + "InstantCharacterGenerate": "InstantCharacter Generate", +} +WEB_DIRECTORY = "./web" + +__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', "WEB_DIRECTORY"] \ No newline at end of file diff --git a/nodes/comfy_nodes.py b/nodes/comfy_nodes.py new file mode 100644 index 0000000..e08cc23 --- /dev/null +++ b/nodes/comfy_nodes.py @@ -0,0 +1,120 @@ +import os +import sys +import torch +import folder_paths +from PIL import Image +import numpy as np + + +# Add the parent directory to the Python path so we can import from easycontrol +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..'))) + +from pipeline import InstantCharacterFluxPipeline +from huggingface_hub import login + + +class InstantCharacterLoadModel: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "hf_token": ("STRING", {"default": "", "multiline": True}), + "ip_adapter_name": (folder_paths.get_filename_list("ipadapter")), + "cpu_offload": ("BOOLEAN", {"default": True}) + } + } + + RETURN_TYPES = ("INSTANTCHAR_PIPE",) + FUNCTION = "load_model" + CATEGORY = "InstantCharacter" + + def load_model(self, hf_token, ip_adapter_name, cpu_offload): + login(token=hf_token) + base_model = "black-forest-labs/FLUX.1-dev" + image_encoder_path = "google/siglip-so400m-patch14-384" + image_encoder_2_path = "facebook/dinov2-giant" + cache_dir = folder_paths.get_folder_paths("diffusers")[0] + device = "cuda" if torch.cuda.is_available() else "cpu" + ip_adapter_path = folder_paths.get_full_path("ipadapter", ip_adapter_name) + pipe = InstantCharacterFluxPipeline.from_pretrained( + base_model, + torch_dtype=torch.bfloat16, + cache_dir=cache_dir, + ) + if cpu_offload: + pipe.enable_sequential_cpu_offload() + else: + pipe.to(device) + + pipe.init_adapter( + image_encoder_path=image_encoder_path, + image_encoder_2_path=image_encoder_2_path, + subject_ipadapter_cfg=dict( + subject_ip_adapter_path=ip_adapter_path, + nb_token=1024 + ), + ) + + return (pipe,) + + +class InstantCharacterGenerate: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "pipe": ("INSTANTCHAR_PIPE",), + "prompt": ("STRING", {"multiline": True}), + "height": ("INT", {"default": 768, "min": 256, "max": 2048, "step": 64}), + "width": ("INT", {"default": 1024, "min": 256, "max": 2048, "step": 64}), + "guidance_scale": ("FLOAT", {"default": 3.5, "min": 0.0, "max": 10.0, "step": 0.1}), + "num_inference_steps": ("INT", {"default": 28, "min": 1, "max": 100, "step": 1}), + "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), + "subject_scale": ("FLOAT", {"default": 0.9, "min": 0.0, "max": 2.0, "step": 0.1}), + }, + "optional": { + "subject_image": ("IMAGE",), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "generate" + CATEGORY = "InstantCharacter" + + def generate(self, pipe, prompt, height, width, guidance_scale, + num_inference_steps, seed, subject_scale, subject_image=None): + + # Convert subject image from tensor to PIL if provided + subject_image_pil = None + if subject_image is not None: + if isinstance(subject_image, torch.Tensor): + if subject_image.dim() == 4: # [batch, height, width, channels] + img = subject_image[0].cpu().numpy() + else: # [height, width, channels] + img = subject_image.cpu().numpy() + subject_image_pil = Image.fromarray((img * 255).astype(np.uint8)) + elif isinstance(subject_image, np.ndarray): + subject_image_pil = Image.fromarray((subject_image * 255).astype(np.uint8)) + + # Generate image + output = pipe( + prompt=prompt, + height=height, + width=width, + guidance_scale=guidance_scale, + num_inference_steps=num_inference_steps, + generator=torch.Generator("cpu").manual_seed(seed), + subject_image=subject_image_pil, + subject_scale=subject_scale, + ) + + # Convert PIL image to tensor format + image = np.array(output.images[0]) / 255.0 + image = torch.from_numpy(image).float() + + # Add batch dimension if needed + if image.dim() == 3: + image = image.unsqueeze(0) + + return (image,) + diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..324bb9c --- /dev/null +++ b/requirements.txt @@ -0,0 +1,10 @@ +diffusers==0.32.2 +easydict +einops +peft +pillow +protobuf +requests +safetensors +sentencepiece +transformers \ No newline at end of file