init: Initial repo, release 0.0.1

This commit is contained in:
JettHu
2024-04-19 14:47:59 +08:00
commit f26e9859d6
10 changed files with 860 additions and 0 deletions
+106
View File
@@ -0,0 +1,106 @@
# ComfyUI-ELLA
<div align="center">
<img src="./assets/ELLA-Diffusion.jpg" width="30%" > <br/>
<a href='https://ella-diffusion.github.io/'><img src='https://img.shields.io/badge/Project-Page-green'></a>
<a href='https://arxiv.org/abs/2403.05135'><img src='https://img.shields.io/badge/arXiv-2403.05135-b31b1b.svg'></a>
</div>
[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}
}
```
+1
View File
@@ -0,0 +1 @@
from .ella import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
+119
View File
@@ -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)
Binary file not shown.

After

Width:  |  Height:  |  Size: 83 KiB

+290
View File
@@ -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",
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 901 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 719 KiB

+296
View File
@@ -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)
+45
View File
@@ -0,0 +1,45 @@
[tool.poetry]
name = "comfyui-ella"
version = "0.1.0"
description = "ELLA plugin for ComfyUI"
authors = ["jetthu <jett.hux@gmail.com>"]
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"
+3
View File
@@ -0,0 +1,3 @@
torch
safetensors
transformers