commit f26e9859d6ea39faadd3211ce690c5ea03a198d4 Author: JettHu Date: Fri Apr 19 14:26:43 2024 +0800 init: Initial repo, release 0.0.1 diff --git a/README.md b/README.md new file mode 100644 index 0000000..479fe73 --- /dev/null +++ b/README.md @@ -0,0 +1,106 @@ +# ComfyUI-ELLA + +
+
+ + +
+ + +[ComfyUI](https://github.com/comfyanonymous/ComfyUI) implementation for [ELLA](https://github.com/TencentQQGYLab/ELLA). + +## :star2: Changelog + +- **[2024.4.19]** Initial repo + +## :books: Example workflows + +The [examples directory](./examples/) has workflow examples. You can directly load these images as workflow into ComfyUI for use. + +![workflow_example](./examples/workflow_example.png) + +:tada: It works with controlnet! And [EMMA](https://github.com/TencentQQGYLab/ELLA/issues/15) is working in progress. + +![controlnet_workflow_example](./examples/controlnet_workflow_example.png) + +## :green_book: Install + +Download or git clone this repository inside ComfyUI/custom_nodes/ directory. `ComfyUI-ELLA` requires the latest version of ComfyUI. If something doesn't work be sure to upgrade. + +```bash +cd ComfyUI/custom_nodes +git clone https://github.com/TencentQQGYLab/ComfyUI-ELLA +``` + +Next install dependencies. + +```bash +cd ComfyUI-ELLA +pip install -r requirements.txt +``` + +## :orange_book: Models + +These models must be placed in the corresponding directories under models. + +Remember you can also use any custom location setting an `ella` & `ella_encoder` entry in the `extra_model_paths.yaml` file. + +- `ComfyUI/models/ella`, create it if not present. + - Place [ELLA Models](https://huggingface.co/QQGYLab/ELLA) here +- `ComfyUI/models/ella_encoder`, create it if not present. + - Place [t5 model](https://huggingface.co/google/flan-t5-xl) here, it should be a folder of transfomers structure with config.json + +In summary, you should have the following model directory structure: + +```bash +ComfyUI/models/ella/ +└── ella-sd1.5-tsc-t5xl.safetensors + +# It may be a little different because I only include the text_encoder part here +ComfyUI/models/ella_encoder/ +└── models--google--flan-t5-xl + ├── config.json + ├── model.safetensors + ├── special_tokens_map.json + ├── spiece.model + ├── tokenizer_config.json + └── tokenizer.json +``` + + +## :book: Usage (:construction: Work in progress) + +#### Load ELLA Model + +#### Apply ELLA + +#### Load T5 TextEncoder #ELLA + +#### T5 Text Encode #ELLA + +#### ELLA Combine Embeds + +#### Convert Condition to ELLA Embeds + + +## :memo: TODO + +- [ ] Support prompt weighting + +## :yum: Thanks + +- ComfyUI: https://github.com/comfyanonymous/ComfyUI +- Diffusers (borrowed timestep modules): https://github.com/huggingface/diffusers + +## :wink: Citation + +``` +@misc{hu2024ella, + title={ELLA: Equip Diffusion Models with LLM for Enhanced Semantic Alignment}, + author={Xiwei Hu and Rui Wang and Yixiao Fang and Bin Fu and Pei Cheng and Gang Yu}, + year={2024}, + eprint={2403.05135}, + archivePrefix={arXiv}, + primaryClass={cs.CV} +} +``` diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..132e80c --- /dev/null +++ b/__init__.py @@ -0,0 +1 @@ +from .ella import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS diff --git a/activations.py b/activations.py new file mode 100644 index 0000000..c4f0460 --- /dev/null +++ b/activations.py @@ -0,0 +1,119 @@ +# coding=utf-8 +# Copyright 2024 HuggingFace Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import torch +import torch.nn.functional as F +from torch import nn + +ACTIVATION_FUNCTIONS = { + "swish": nn.SiLU(), + "silu": nn.SiLU(), + "mish": nn.Mish(), + "gelu": nn.GELU(), + "relu": nn.ReLU(), +} + + +def get_activation(act_fn: str) -> nn.Module: + """Helper function to get activation function from string. + + Args: + act_fn (str): Name of activation function. + + Returns: + nn.Module: Activation function. + """ + + act_fn = act_fn.lower() + if act_fn in ACTIVATION_FUNCTIONS: + return ACTIVATION_FUNCTIONS[act_fn] + else: + raise ValueError(f"Unsupported activation function: {act_fn}") + + +class GELU(nn.Module): + r""" + GELU activation function with tanh approximation support with `approximate="tanh"`. + + Parameters: + dim_in (`int`): The number of channels in the input. + dim_out (`int`): The number of channels in the output. + approximate (`str`, *optional*, defaults to `"none"`): If `"tanh"`, use tanh approximation. + bias (`bool`, defaults to True): Whether to use a bias in the linear layer. + """ + + def __init__(self, dim_in: int, dim_out: int, approximate: str = "none", bias: bool = True): + super().__init__() + self.proj = nn.Linear(dim_in, dim_out, bias=bias) + self.approximate = approximate + + def gelu(self, gate: torch.Tensor) -> torch.Tensor: + if gate.device.type != "mps": + return F.gelu(gate, approximate=self.approximate) + # mps: gelu is not implemented for float16 + return F.gelu(gate.to(dtype=torch.float32), approximate=self.approximate).to(dtype=gate.dtype) + + def forward(self, hidden_states): + hidden_states = self.proj(hidden_states) + hidden_states = self.gelu(hidden_states) + return hidden_states + + +class GEGLU(nn.Module): + r""" + A [variant](https://arxiv.org/abs/2002.05202) of the gated linear unit activation function. + + Parameters: + dim_in (`int`): The number of channels in the input. + dim_out (`int`): The number of channels in the output. + bias (`bool`, defaults to True): Whether to use a bias in the linear layer. + """ + + def __init__(self, dim_in: int, dim_out: int, bias: bool = True): + super().__init__() + self.proj = nn.Linear(dim_in, dim_out * 2, bias=bias) + + def gelu(self, gate: torch.Tensor) -> torch.Tensor: + if gate.device.type != "mps": + return F.gelu(gate) + # mps: gelu is not implemented for float16 + return F.gelu(gate.to(dtype=torch.float32)).to(dtype=gate.dtype) + + def forward(self, hidden_states, *args, **kwargs): + if len(args) > 0 or kwargs.get("scale", None) is not None: + raise ValueError("`scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`.") + + hidden_states, gate = self.proj(hidden_states).chunk(2, dim=-1) + return hidden_states * self.gelu(gate) + + +class ApproximateGELU(nn.Module): + r""" + The approximate form of the Gaussian Error Linear Unit (GELU). For more details, see section 2 of this + [paper](https://arxiv.org/abs/1606.08415). + + Parameters: + dim_in (`int`): The number of channels in the input. + dim_out (`int`): The number of channels in the output. + bias (`bool`, defaults to True): Whether to use a bias in the linear layer. + """ + + def __init__(self, dim_in: int, dim_out: int, bias: bool = True): + super().__init__() + self.proj = nn.Linear(dim_in, dim_out, bias=bias) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = self.proj(x) + return x * torch.sigmoid(1.702 * x) diff --git a/assets/ELLA-Diffusion.jpg b/assets/ELLA-Diffusion.jpg new file mode 100644 index 0000000..ce5af47 Binary files /dev/null and b/assets/ELLA-Diffusion.jpg differ diff --git a/ella.py b/ella.py new file mode 100644 index 0000000..e17d99f --- /dev/null +++ b/ella.py @@ -0,0 +1,290 @@ +import os + +import folder_paths +import torch +from comfy import model_management +from safetensors.torch import load_model + +from .model import ELLA, T5TextEmbedder + +ELLA_EMBEDS_TYPE = "ELLA_EMBEDS" +ELLA_EMBEDS_PREFIX = "ella_" +ELLA_EMBEDS_PREFIX_LEN = len(ELLA_EMBEDS_PREFIX) + + +class EllaProxyUNet: + def __init__(self, ella, model_sampling, positive, negative) -> None: + self.ella = ella + self.model_sampling = model_sampling + if positive.keys() != negative.keys(): + raise ValueError("positive and negative embeds types must match") + self.embeds = [positive, negative] + + self.dtype = model_management.text_encoder_dtype() + self.ella.to(self.dtype) + for i in range(len(self.embeds)): + for k in self.embeds[i]: + self.embeds[i][k].to(device=self.load_device, dtype=self.dtype) + + @property + def load_device(self): + return model_management.text_encoder_device() + + @property + def offload_device(self): + return model_management.text_encoder_offload_device() + + def prepare_conds(self): + self.ella.to(self.load_device) + cond = self.ella(torch.Tensor([999]).to(torch.int64), **self.embeds[0]) + uncond = self.ella(torch.Tensor([999]).to(torch.int64), **self.embeds[1]) + self.ella.to(self.offload_device) + return cond, uncond + + def __call__(self, apply_model, kwargs: dict): + input_x = kwargs["input"] + timestep_ = kwargs["timestep"] + c = kwargs["c"] + cond_or_uncond = kwargs["cond_or_uncond"] # [0|1] + + time_aware_encoder_hidden_states = [] + self.ella.to(device=self.load_device) + for i in cond_or_uncond: + h = self.ella( + self.model_sampling.timestep(timestep_[i]), + **self.embeds[i], + ) + time_aware_encoder_hidden_states.append(h) + self.ella.to(self.offload_device) + + c["c_crossattn"] = torch.cat(time_aware_encoder_hidden_states, dim=0) + + return apply_model(input_x, timestep_, **c) + + +""" +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + Apply Nodes +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ +""" + + +class EllaApply: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": ("MODEL",), + "ella": ("ELLA",), + "positive": (ELLA_EMBEDS_TYPE,), + "negative": (ELLA_EMBEDS_TYPE,), + } + } + + RETURN_NAMES = ("model", "positive", "negative") + RETURN_TYPES = ("MODEL", "CONDITIONING", "CONDITIONING") + FUNCTION = "apply" + CATEGORY = "ella/apply" + + def apply(self, model, ella, positive, negative): + model_clone = model.clone() + model_sampling = model_clone.get_model_object("model_sampling") + + ella_proxy = EllaProxyUNet( + ella=ella, + model_sampling=model_sampling, + positive={ + k[ELLA_EMBEDS_PREFIX_LEN:]: v.clone() for k, v in positive.items() if k.startswith(ELLA_EMBEDS_PREFIX) + }, + negative={ + k[ELLA_EMBEDS_PREFIX_LEN:]: v.clone() for k, v in negative.items() if k.startswith(ELLA_EMBEDS_PREFIX) + }, + ) + + model_clone.set_model_unet_function_wrapper(ella_proxy) + # No matter how many tokens are text features, the ella output must be 64 tokens. + _cond, _uncond = ella_proxy.prepare_conds() + cond = [_cond, {k: v for k, v in positive.items() if not k.startswith(ELLA_EMBEDS_PREFIX)}] + uncond = [_uncond, {k: v for k, v in negative.items() if not k.startswith(ELLA_EMBEDS_PREFIX)}] + + return (model_clone, [cond], [uncond]) + + +""" +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + Encoders +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ +""" + + +class T5TextEncode: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "text": ("STRING", {"multiline": True, "dynamicPrompts": True}), + "text_encoder": ("T5_TEXT_ENCODER",), + } + } + + RETURN_TYPES = (ELLA_EMBEDS_TYPE,) + FUNCTION = "encode" + + CATEGORY = "ella/conditioning" + + def encode(self, text, text_encoder, max_length=None): + # TODO: more offload strategy + text_encoder.to(model_management.text_encoder_device()) + cond = text_encoder(text, max_length=max_length) + text_encoder.to(model_management.text_encoder_offload_device()) + + return ({f"{ELLA_EMBEDS_PREFIX}t5_embeds": cond},) + + +""" +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + Loaders +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ +""" + + +class ELLALoader: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "name": (folder_paths.get_filename_list("ella"),), + } + } + + RETURN_TYPES = ("ELLA",) + FUNCTION = "load" + CATEGORY = "ella/loaders" + + def load(self, name: str, **kwargs): + ella_file = folder_paths.get_full_path("ella", name) + # TODO: expose more ELLA init params or takes from ckpt + ella = ELLA() + load_model(ella, ella_file, strict=True) # type: ignore + return (ella,) + + +class T5TextEncoderLoader: + @classmethod + def INPUT_TYPES(cls): + paths = [] + for search_path in folder_paths.get_folder_paths("ella_encoder"): + if os.path.exists(search_path): + for root, _, files in os.walk(search_path, followlinks=True): + if "config.json" in files: + paths.append(os.path.relpath(root, start=search_path)) + return { + "required": { + "name": (paths,), + "max_length": ("INT", {"default": 0, "min": 0, "max": 128, "step": 16}), + "dtype": (["auto", "FP32", "FP16"],), + } + } + + RETURN_TYPES = ("T5_TEXT_ENCODER",) + FUNCTION = "load" + CATEGORY = "ella/loaders" + + def load(self, name: str, max_length: int = 0, dtype="auto", **kwargs): + t5_file = folder_paths.get_full_path("ella_encoder", name) + # "flexible_token_length" trick: Set `max_length=None` eliminating any text token padding or truncation. + # Help improve the quality of generated images corresponding to short captions. + for search_path in folder_paths.get_folder_paths("ella_encoder"): + if os.path.exists(search_path): + path = os.path.join(search_path, name) + if os.path.exists(path): + t5_file = path + break + t5_encoder = T5TextEmbedder(t5_file, max_length=max_length or None) # type: ignore + if dtype == "auto": + dtype = model_management.text_encoder_dtype() + elif dtype == "FP16": + dtype = torch.float16 + else: + dtype = torch.float32 + t5_encoder.to(dtype) # type: ignore + return (t5_encoder,) + + +""" +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + Helper +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ +""" + + +class ConditionToEllaEmbeds: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "cond": ("CONDITIONING",), + } + } + + RETURN_TYPES = (ELLA_EMBEDS_TYPE,) + FUNCTION = "convert" + + CATEGORY = "ella/helper" + + def convert(self, cond): + # only use batch 0 + # CONDITIONING: [[cond, {"pooled_output": pooled}]] + return ({f"{ELLA_EMBEDS_PREFIX}clip_embeds": cond[0][0], **cond[0][1]},) + + +class EllaCombineEmbeds: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "embeds": (ELLA_EMBEDS_TYPE,), + "embeds_add": (ELLA_EMBEDS_TYPE,), + } + } + + RETURN_TYPES = (ELLA_EMBEDS_TYPE,) + FUNCTION = "combine" + + CATEGORY = "ella/helper" + + def combine(self, embeds: dict, embeds_add: dict): + if embeds.keys() & embeds_add.keys(): + print("warning: because there are some same keys, one of them will be overwritten.") + + return ({**embeds, **embeds_add},) + + +""" +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + Register +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ +""" +NODE_CLASS_MAPPINGS = { + # Main Apply Nodes + "EllaApply": EllaApply, + "T5TextEncode #ELLA": T5TextEncode, + # Loaders + "ELLALoader": ELLALoader, + "T5TextEncoderLoader #ELLA": T5TextEncoderLoader, + # Helpers + "EllaCombineEmbeds": EllaCombineEmbeds, + "ConditionToEllaEmbeds": ConditionToEllaEmbeds, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + # Main Apply Nodes + "EllaApply": "Apply ELLA", + "T5TextEncode #ELLA": "T5 Text Encode #ELLA", + # Loaders + "ELLALoader": "Load ELLA Model", + "T5TextEncoderLoader #ELLA": "Load T5 TextEncoder #ELLA", + # Helpers + "EllaCombineEmbeds": "ELLA Combine Embeds", + "ConditionToEllaEmbeds": "Convert Condition to ELLA Embeds", +} diff --git a/examples/controlnet_workflow_example.png b/examples/controlnet_workflow_example.png new file mode 100644 index 0000000..541fed3 Binary files /dev/null and b/examples/controlnet_workflow_example.png differ diff --git a/examples/workflow_example.png b/examples/workflow_example.png new file mode 100644 index 0000000..056a7a3 Binary files /dev/null and b/examples/workflow_example.png differ diff --git a/model.py b/model.py new file mode 100644 index 0000000..fa40777 --- /dev/null +++ b/model.py @@ -0,0 +1,296 @@ +import math +from collections import OrderedDict +from typing import Optional + +import torch +from torch import nn + +from .activations import get_activation + + +class AdaLayerNorm(nn.Module): + def __init__(self, embedding_dim: int, time_embedding_dim: Optional[int] = None): + super().__init__() + + if time_embedding_dim is None: + time_embedding_dim = embedding_dim + + self.silu = nn.SiLU() + self.linear = nn.Linear(time_embedding_dim, 2 * embedding_dim, bias=True) + nn.init.zeros_(self.linear.weight) + nn.init.zeros_(self.linear.bias) + + self.norm = nn.LayerNorm(embedding_dim, elementwise_affine=False, eps=1e-6) + + def forward(self, x: torch.Tensor, timestep_embedding: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + emb = self.linear(self.silu(timestep_embedding)) + shift, scale = emb.view(len(x), 1, -1).chunk(2, dim=-1) + return self.norm(x) * (1 + scale) + shift + + +class SquaredReLU(nn.Module): + def forward(self, x: torch.Tensor): + return torch.square(torch.relu(x)) + + +class PerceiverAttentionBlock(nn.Module): + def __init__(self, d_model: int, n_heads: int, time_embedding_dim: Optional[int] = None): + super().__init__() + self.attn = nn.MultiheadAttention(d_model, n_heads, batch_first=True) + + self.mlp = nn.Sequential( + OrderedDict( + [ + ("c_fc", nn.Linear(d_model, d_model * 4)), + ("sq_relu", SquaredReLU()), + ("c_proj", nn.Linear(d_model * 4, d_model)), + ] + ) + ) + + self.ln_1 = AdaLayerNorm(d_model, time_embedding_dim) + self.ln_2 = AdaLayerNorm(d_model, time_embedding_dim) + self.ln_ff = AdaLayerNorm(d_model, time_embedding_dim) + + def attention(self, q: torch.Tensor, kv: torch.Tensor): + attn_output, attn_output_weights = self.attn(q, kv, kv, need_weights=False) + return attn_output + + def forward( + self, + x: torch.Tensor, + latents: torch.Tensor, + timestep_embedding: Optional[torch.Tensor] = None, + ): + normed_latents = self.ln_1(latents, timestep_embedding) + latents = latents + self.attention( + q=normed_latents, + kv=torch.cat([normed_latents, self.ln_2(x, timestep_embedding)], dim=1), + ) + return latents + self.mlp(self.ln_ff(latents, timestep_embedding)) + + +class PerceiverResampler(nn.Module): + def __init__( + self, + width: int = 768, + layers: int = 6, + heads: int = 8, + num_latents: int = 64, + output_dim=None, + input_dim=None, + time_embedding_dim: Optional[int] = None, + ): + super().__init__() + self.output_dim = output_dim + self.input_dim = input_dim + self.latents = nn.Parameter(width**-0.5 * torch.randn(num_latents, width)) + self.time_aware_linear = nn.Linear(time_embedding_dim or width, width, bias=True) + + if self.input_dim is not None: + self.proj_in = nn.Linear(input_dim, width) # type: ignore + + self.perceiver_blocks = nn.Sequential( + *[PerceiverAttentionBlock(width, heads, time_embedding_dim=time_embedding_dim) for _ in range(layers)] + ) + + if self.output_dim is not None: + self.proj_out = nn.Sequential(nn.Linear(width, output_dim), nn.LayerNorm(output_dim)) # type: ignore + + def forward(self, x: torch.Tensor, timestep_embedding: torch.Tensor = None): # type: ignore + learnable_latents = self.latents.unsqueeze(dim=0).repeat(len(x), 1, 1) + latents = learnable_latents + self.time_aware_linear(torch.nn.functional.silu(timestep_embedding)) + if self.input_dim is not None: + x = self.proj_in(x) + for p_block in self.perceiver_blocks: + latents = p_block(x, latents, timestep_embedding=timestep_embedding) + + if self.output_dim is not None: + latents = self.proj_out(latents) + + return latents + + +class T5TextEmbedder(nn.Module): + def __init__(self, pretrained_path="google/flan-t5-xl", max_length=None): + super().__init__() + # TODO: make it naive instead of transformers + from transformers import T5EncoderModel, T5Tokenizer + + self.model = T5EncoderModel.from_pretrained(pretrained_path) + self.tokenizer = T5Tokenizer.from_pretrained(pretrained_path) + self.max_length = max_length + + def forward(self, caption, text_input_ids=None, attention_mask=None, max_length=None, **kwargs): + if max_length is None: + max_length = self.max_length + + if text_input_ids is None or attention_mask is None: + if max_length is not None: + text_inputs = self.tokenizer( + caption, + return_tensors="pt", + add_special_tokens=True, + max_length=max_length, + padding="max_length", + truncation=True, + ) + else: + text_inputs = self.tokenizer(caption, return_tensors="pt", add_special_tokens=True) + text_input_ids = text_inputs.input_ids + attention_mask = text_inputs.attention_mask + text_input_ids = text_input_ids.to(self.model.device) # type: ignore + attention_mask = attention_mask.to(self.model.device) # type: ignore + outputs = self.model(text_input_ids, attention_mask=attention_mask) # type: ignore + + return outputs.last_hidden_state + + +class TimestepEmbedding(nn.Module): + def __init__( + self, + in_channels: int, + time_embed_dim: int, + act_fn: str = "silu", + out_dim: Optional[int] = None, + post_act_fn: Optional[str] = None, + cond_proj_dim=None, + sample_proj_bias=True, + ): + super().__init__() + linear_cls = nn.Linear + + self.linear_1 = linear_cls(in_channels, time_embed_dim, sample_proj_bias) + + if cond_proj_dim is not None: + self.cond_proj = nn.Linear(cond_proj_dim, in_channels, bias=False) + else: + self.cond_proj = None + + self.act = get_activation(act_fn) + + time_embed_dim_out = out_dim if out_dim is not None else time_embed_dim + self.linear_2 = linear_cls(time_embed_dim, time_embed_dim_out, sample_proj_bias) + + if post_act_fn is None: + self.post_act = None + else: + self.post_act = get_activation(post_act_fn) + + def forward(self, sample, condition=None): + if condition is not None: + sample = sample + self.cond_proj(condition) # type: ignore + sample = self.linear_1(sample) + + if self.act is not None: + sample = self.act(sample) + + sample = self.linear_2(sample) + + if self.post_act is not None: + sample = self.post_act(sample) + return sample + + +def get_timestep_embedding( + timesteps: torch.Tensor, + embedding_dim: int, + flip_sin_to_cos: bool = False, + downscale_freq_shift: float = 1, + scale: float = 1, + max_period: int = 10000, +): + """ + This matches the implementation in Denoising Diffusion Probabilistic Models: Create sinusoidal timestep embeddings. + + :param timesteps: a 1-D Tensor of N indices, one per batch element. + These may be fractional. + :param embedding_dim: the dimension of the output. :param max_period: controls the minimum frequency of the + embeddings. :return: an [N x dim] Tensor of positional embeddings. + """ + assert len(timesteps.shape) == 1, "Timesteps should be a 1d-array" + + half_dim = embedding_dim // 2 + exponent = -math.log(max_period) * torch.arange( + start=0, end=half_dim, dtype=torch.float32, device=timesteps.device + ) + exponent = exponent / (half_dim - downscale_freq_shift) + + emb = torch.exp(exponent) + emb = timesteps[:, None].float() * emb[None, :] + + # scale embeddings + emb = scale * emb + + # concat sine and cosine embeddings + emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1) + + # flip sine and cosine embeddings + if flip_sin_to_cos: + emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1) + + # zero pad + if embedding_dim % 2 == 1: + emb = torch.nn.functional.pad(emb, (0, 1, 0, 0)) + return emb + + +class Timesteps(nn.Module): + def __init__(self, num_channels: int, flip_sin_to_cos: bool, downscale_freq_shift: float): + super().__init__() + self.num_channels = num_channels + self.flip_sin_to_cos = flip_sin_to_cos + self.downscale_freq_shift = downscale_freq_shift + + def forward(self, timesteps): + return get_timestep_embedding( + timesteps, + self.num_channels, + flip_sin_to_cos=self.flip_sin_to_cos, + downscale_freq_shift=self.downscale_freq_shift, + ) + + +class ELLA(nn.Module): + def __init__( + self, + time_channel=320, + time_embed_dim=768, + act_fn: str = "silu", + out_dim: Optional[int] = None, + width=768, + layers=6, + heads=8, + num_latents=64, + input_dim=2048, + ): + super().__init__() + # TODO: make it naive instead of diffusers.models.embeddings.[Timesteps, TimestepEmbedding] + + self.position = Timesteps(time_channel, flip_sin_to_cos=True, downscale_freq_shift=0) + self.time_embedding = TimestepEmbedding( + in_channels=time_channel, + time_embed_dim=time_embed_dim, + act_fn=act_fn, + out_dim=out_dim, # type: ignore + ) + + self.connector = PerceiverResampler( + width=width, + layers=layers, + heads=heads, + num_latents=num_latents, + input_dim=input_dim, + time_embedding_dim=time_embed_dim, + ) + + def forward(self, timesteps: torch.Tensor, t5_embeds: torch.Tensor, **kwargs): + device = t5_embeds.device + dtype = t5_embeds.dtype + + ori_time_feature = self.position(timesteps.view(-1)).to(device, dtype=dtype) + ori_time_feature = ori_time_feature.unsqueeze(dim=1) if ori_time_feature.ndim == 2 else ori_time_feature + ori_time_feature = ori_time_feature.expand(len(t5_embeds), -1, -1) + time_embedding = self.time_embedding(ori_time_feature) + + return self.connector(t5_embeds, timestep_embedding=time_embedding) diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..aa6284c --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,45 @@ +[tool.poetry] +name = "comfyui-ella" +version = "0.1.0" +description = "ELLA plugin for ComfyUI" +authors = ["jetthu "] +license = "GPL-3.0-only" +readme = "README.md" +packages = [{ include = "*.py" }] + + +[tool.poetry.dependencies] +python = ">=3.6" + +[tool.ruff] +line-length = 119 + +# A list of file patterns to omit from linting, in addition to those specified by exclude. +extend-exclude = ["__pycache__", "*.pyc", "*.egg-info", ".cache"] + +select = ["E", "F", "W", "C90", "I", "UP", "B", "C4", "RET", "RUF", "SIM"] + + +ignore = [ + "UP006", # UP006: Use list instead of typing.List for type annotations + "UP007", # UP007: Use X | Y for type annotations + "UP009", + "UP035", + "UP038", + "E402", +] + +[tool.ruff.per-file-ignores] +# F401: unused-import +"__init__.py" = ["F401"] + +[tool.isort] +profile = "black" + +[tool.black] +line-length = 119 +skip-string-normalization = true + +[build-system] +requires = ["poetry-core"] +build-backend = "poetry.core.masonry.api" diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..c78a037 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,3 @@ +torch +safetensors +transformers