init: Initial repo, release 0.0.1
This commit is contained in:
@@ -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.
|
||||
|
||||

|
||||
|
||||
:tada: It works with controlnet! And [EMMA](https://github.com/TencentQQGYLab/ELLA/issues/15) is working in progress.
|
||||
|
||||

|
||||
|
||||
## :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}
|
||||
}
|
||||
```
|
||||
@@ -0,0 +1 @@
|
||||
from .ella import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
+119
@@ -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 |
@@ -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 |
@@ -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)
|
||||
@@ -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"
|
||||
@@ -0,0 +1,3 @@
|
||||
torch
|
||||
safetensors
|
||||
transformers
|
||||
Reference in New Issue
Block a user