Initial commit

This commit is contained in:
Kijai
2024-04-09 19:32:57 +03:00
parent 6bae7dba83
commit 237e0e3e40
21 changed files with 148992 additions and 0 deletions
+6
View File
@@ -0,0 +1,6 @@
__pycache__/
checkpoints/
*.py[cod]
*$py.class
*.egg-info
.pytest_cache
+201
View File
@@ -0,0 +1,201 @@
Apache License
Version 2.0, January 2004
http://www.apache.org/licenses/
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
1. Definitions.
"License" shall mean the terms and conditions for use, reproduction,
and distribution as defined by Sections 1 through 9 of this document.
"Licensor" shall mean the copyright owner or entity authorized by
the copyright owner that is granting the License.
"Legal Entity" shall mean the union of the acting entity and all
other entities that control, are controlled by, or are under common
control with that entity. For the purposes of this definition,
"control" means (i) the power, direct or indirect, to cause the
direction or management of such entity, whether by contract or
otherwise, or (ii) ownership of fifty percent (50%) or more of the
outstanding shares, or (iii) beneficial ownership of such entity.
"You" (or "Your") shall mean an individual or Legal Entity
exercising permissions granted by this License.
"Source" form shall mean the preferred form for making modifications,
including but not limited to software source code, documentation
source, and configuration files.
"Object" form shall mean any form resulting from mechanical
transformation or translation of a Source form, including but
not limited to compiled object code, generated documentation,
and conversions to other media types.
"Work" shall mean the work of authorship, whether in Source or
Object form, made available under the License, as indicated by a
copyright notice that is included in or attached to the work
(an example is provided in the Appendix below).
"Derivative Works" shall mean any work, whether in Source or Object
form, that is based on (or derived from) the Work and for which the
editorial revisions, annotations, elaborations, or other modifications
represent, as a whole, an original work of authorship. For the purposes
of this License, Derivative Works shall not include works that remain
separable from, or merely link (or bind by name) to the interfaces of,
the Work and Derivative Works thereof.
"Contribution" shall mean any work of authorship, including
the original version of the Work and any modifications or additions
to that Work or Derivative Works thereof, that is intentionally
submitted to Licensor for inclusion in the Work by the copyright owner
or by an individual or Legal Entity authorized to submit on behalf of
the copyright owner. For the purposes of this definition, "submitted"
means any form of electronic, verbal, or written communication sent
to the Licensor or its representatives, including but not limited to
communication on electronic mailing lists, source code control systems,
and issue tracking systems that are managed by, or on behalf of, the
Licensor for the purpose of discussing and improving the Work, but
excluding communication that is conspicuously marked or otherwise
designated in writing by the copyright owner as "Not a Contribution."
"Contributor" shall mean Licensor and any individual or Legal Entity
on behalf of whom a Contribution has been received by Licensor and
subsequently incorporated within the Work.
2. Grant of Copyright License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
copyright license to reproduce, prepare Derivative Works of,
publicly display, publicly perform, sublicense, and distribute the
Work and such Derivative Works in Source or Object form.
3. Grant of Patent License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
(except as stated in this section) patent license to make, have made,
use, offer to sell, sell, import, and otherwise transfer the Work,
where such license applies only to those patent claims licensable
by such Contributor that are necessarily infringed by their
Contribution(s) alone or by combination of their Contribution(s)
with the Work to which such Contribution(s) was submitted. If You
institute patent litigation against any entity (including a
cross-claim or counterclaim in a lawsuit) alleging that the Work
or a Contribution incorporated within the Work constitutes direct
or contributory patent infringement, then any patent licenses
granted to You under this License for that Work shall terminate
as of the date such litigation is filed.
4. Redistribution. You may reproduce and distribute copies of the
Work or Derivative Works thereof in any medium, with or without
modifications, and in Source or Object form, provided that You
meet the following conditions:
(a) You must give any other recipients of the Work or
Derivative Works a copy of this License; and
(b) You must cause any modified files to carry prominent notices
stating that You changed the files; and
(c) You must retain, in the Source form of any Derivative Works
that You distribute, all copyright, patent, trademark, and
attribution notices from the Source form of the Work,
excluding those notices that do not pertain to any part of
the Derivative Works; and
(d) If the Work includes a "NOTICE" text file as part of its
distribution, then any Derivative Works that You distribute must
include a readable copy of the attribution notices contained
within such NOTICE file, excluding those notices that do not
pertain to any part of the Derivative Works, in at least one
of the following places: within a NOTICE text file distributed
as part of the Derivative Works; within the Source form or
documentation, if provided along with the Derivative Works; or,
within a display generated by the Derivative Works, if and
wherever such third-party notices normally appear. The contents
of the NOTICE file are for informational purposes only and
do not modify the License. You may add Your own attribution
notices within Derivative Works that You distribute, alongside
or as an addendum to the NOTICE text from the Work, provided
that such additional attribution notices cannot be construed
as modifying the License.
You may add Your own copyright statement to Your modifications and
may provide additional or different license terms and conditions
for use, reproduction, or distribution of Your modifications, or
for any such Derivative Works as a whole, provided Your use,
reproduction, and distribution of the Work otherwise complies with
the conditions stated in this License.
5. Submission of Contributions. Unless You explicitly state otherwise,
any Contribution intentionally submitted for inclusion in the Work
by You to the Licensor shall be under the terms and conditions of
this License, without any additional terms or conditions.
Notwithstanding the above, nothing herein shall supersede or modify
the terms of any separate license agreement you may have executed
with Licensor regarding such Contributions.
6. Trademarks. This License does not grant permission to use the trade
names, trademarks, service marks, or product names of the Licensor,
except as required for reasonable and customary use in describing the
origin of the Work and reproducing the content of the NOTICE file.
7. Disclaimer of Warranty. Unless required by applicable law or
agreed to in writing, Licensor provides the Work (and each
Contributor provides its Contributions) on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
implied, including, without limitation, any warranties or conditions
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
PARTICULAR PURPOSE. You are solely responsible for determining the
appropriateness of using or redistributing the Work and assume any
risks associated with Your exercise of permissions under this License.
8. Limitation of Liability. In no event and under no legal theory,
whether in tort (including negligence), contract, or otherwise,
unless required by applicable law (such as deliberate and grossly
negligent acts) or agreed to in writing, shall any Contributor be
liable to You for damages, including any direct, indirect, special,
incidental, or consequential damages of any character arising as a
result of this License or out of the use or inability to use the
Work (including but not limited to damages for loss of goodwill,
work stoppage, computer failure or malfunction, or any and all
other commercial damages or losses), even if such Contributor
has been advised of the possibility of such damages.
9. Accepting Warranty or Additional Liability. While redistributing
the Work or Derivative Works thereof, You may choose to offer,
and charge a fee for, acceptance of support, warranty, indemnity,
or other liability obligations and/or rights consistent with this
License. However, in accepting such obligations, You may act only
on Your own behalf and on Your sole responsibility, not on behalf
of any other Contributor, and only if You agree to indemnify,
defend, and hold each Contributor harmless for any liability
incurred by, or claims asserted against, such Contributor by reason
of your accepting any such warranty or additional liability.
END OF TERMS AND CONDITIONS
APPENDIX: How to apply the Apache License to your work.
To apply the Apache License to your work, attach the following
boilerplate notice, with the fields enclosed by brackets "[]"
replaced with your own identifying information. (Don't include
the brackets!) The text should be enclosed in the appropriate
comment syntax for the file format. We also recommend that a
file or class name and description of purpose be included on the
same "printed page" as the copyright notice for easier
identification within third-party archives.
Copyright [yyyy] [name of copyright owner]
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.
+103
View File
@@ -0,0 +1,103 @@
# ELLA: Equip Diffusion Models with LLM for Enhanced Semantic Alignment
<div align="center">
<span class="author-block">
<a href="https://openreview.net/profile?id=~Xiwei_Hu1">Xiwei Hu*</a>,
</span>
<span class="author-block">
<a href="https://wrong.wang/">Rui Wang*</a>,
</span>
<span class="author-block">
<a href="https://openreview.net/profile?id=~Yixiao_Fang1">Yixiao Fang*</a>,
</span>
<span class="author-block">
<a href="https://openreview.net/profile?id=~BIN_FU2">Bin Fu*</a>,
</span>
<span class="author-block">
<a href="https://openreview.net/profile?id=~Pei_Cheng1">Pei Cheng</a>,
</span>
<span class="author-block">
<a href="https://www.skicyyu.org/">Gang Yu&#10022</a>
</span>
<p>
* Equal contributions, &#10022 Corresponding Author
</p>
<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>
Official code of "ELLA: Equip Diffusion Models with LLM for Enhanced Semantic Alignment".
<p>
</p>
<div align="center">
<img src="./assets/teaser_3img.png" width="100%">
<img src="./assets/teaser1_raccoon.png" width="100%">
</div>
## 🌟 Changelog
- **[2024.4.9]** 🔥🔥🔥 Release [ELLA-SD1.5](https://huggingface.co/QQGYLab/ELLA/blob/main/ella-sd1.5-tsc-t5xl.safetensors) Checkpoint! Welcome to try!
- **[2024.3.11]** 🔥 Release DPG-Bench! Welcome to try!
- **[2024.3.7]** Initial update
## Inference
### ELLA-SD1.5
```bash
# get ELLA-SD1.5 at https://huggingface.co/QQGYLab/ELLA/blob/main/ella-sd1.5-tsc-t5xl.safetensors
# comparing ella-sd1.5 and sd1.5
# will generate images at `./assets/ella-inference-examples`
python3 inference.py test --save_folder ./assets/ella-inference-examples --ella_path /path/to/ella-sd1.5-tsc-t5xl.safetensors
# build a demo for ella-sd1.5
GRADIO_SERVER_NAME=0.0.0.0 GRADIO_SERVER_PORT=8082 python3 ./inference.py demo /path/to/ella-sd1.5-tsc-t5xl.safetensors
```
## 📊 DPG-Bench
The guideline of DPG-Bench:
1. Generate your images according to our [prompts](./dpg_bench/prompts/).
It is recommended to generate 4 images per prompt and grid them to 2x2 format. **Please Make sure your generated image's filename is the same with the prompt's filename.**
2. Run the following command to conduct evaluation.
```bash
bash dpg_bench/dist_eval.sh $YOUR_IMAGE_PATH $RESOLUTION
```
Thanks to the excellent work of [DSG](https://github.com/j-min/DSG) sincerely, we follow their instructions to generate questions and answers of DPG-Bench.
## 📝 TODO
- [ ] add huggingface demo link
- [x] release checkpoint
- [x] release inference code
- [x] release DPG-Bench
## 💡 Others
We have also found [LaVi-Bridge](https://arxiv.org/abs/2403.07860), another independent but similar work completed almost concurrently, which offers additional insights not covered by ELLA. The difference between ELLA and LaVi-Bridge can be found in [issue 13](https://github.com/ELLA-Diffusion/ELLA/issues/13). We are delighted to welcome other researchers and community users to promote the development of this field.
## 😉 Citation
If you find **ELLA** useful for your research and applications, please cite us using this BibTeX:
```
@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}
}
```
+3
View File
@@ -0,0 +1,3 @@
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
Binary file not shown.

After

Width:  |  Height:  |  Size: 83 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 4.2 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 6.1 MiB

+25
View File
@@ -0,0 +1,25 @@
{
"_name_or_path": "openai/clip-vit-large-patch14",
"architectures": [
"CLIPTextModel"
],
"attention_dropout": 0.0,
"bos_token_id": 0,
"dropout": 0.0,
"eos_token_id": 2,
"hidden_act": "quick_gelu",
"hidden_size": 768,
"initializer_factor": 1.0,
"initializer_range": 0.02,
"intermediate_size": 3072,
"layer_norm_eps": 1e-05,
"max_position_embeddings": 77,
"model_type": "clip_text_model",
"num_attention_heads": 12,
"num_hidden_layers": 12,
"pad_token_id": 1,
"projection_dim": 768,
"torch_dtype": "float32",
"transformers_version": "4.22.0.dev0",
"vocab_size": 49408
}
+171
View File
@@ -0,0 +1,171 @@
{
"_name_or_path": "clip-vit-large-patch14/",
"architectures": [
"CLIPModel"
],
"initializer_factor": 1.0,
"logit_scale_init_value": 2.6592,
"model_type": "clip",
"projection_dim": 768,
"text_config": {
"_name_or_path": "",
"add_cross_attention": false,
"architectures": null,
"attention_dropout": 0.0,
"bad_words_ids": null,
"bos_token_id": 0,
"chunk_size_feed_forward": 0,
"cross_attention_hidden_size": null,
"decoder_start_token_id": null,
"diversity_penalty": 0.0,
"do_sample": false,
"dropout": 0.0,
"early_stopping": false,
"encoder_no_repeat_ngram_size": 0,
"eos_token_id": 2,
"finetuning_task": null,
"forced_bos_token_id": null,
"forced_eos_token_id": null,
"hidden_act": "quick_gelu",
"hidden_size": 768,
"id2label": {
"0": "LABEL_0",
"1": "LABEL_1"
},
"initializer_factor": 1.0,
"initializer_range": 0.02,
"intermediate_size": 3072,
"is_decoder": false,
"is_encoder_decoder": false,
"label2id": {
"LABEL_0": 0,
"LABEL_1": 1
},
"layer_norm_eps": 1e-05,
"length_penalty": 1.0,
"max_length": 20,
"max_position_embeddings": 77,
"min_length": 0,
"model_type": "clip_text_model",
"no_repeat_ngram_size": 0,
"num_attention_heads": 12,
"num_beam_groups": 1,
"num_beams": 1,
"num_hidden_layers": 12,
"num_return_sequences": 1,
"output_attentions": false,
"output_hidden_states": false,
"output_scores": false,
"pad_token_id": 1,
"prefix": null,
"problem_type": null,
"projection_dim" : 768,
"pruned_heads": {},
"remove_invalid_values": false,
"repetition_penalty": 1.0,
"return_dict": true,
"return_dict_in_generate": false,
"sep_token_id": null,
"task_specific_params": null,
"temperature": 1.0,
"tie_encoder_decoder": false,
"tie_word_embeddings": true,
"tokenizer_class": null,
"top_k": 50,
"top_p": 1.0,
"torch_dtype": null,
"torchscript": false,
"transformers_version": "4.16.0.dev0",
"use_bfloat16": false,
"vocab_size": 49408
},
"text_config_dict": {
"hidden_size": 768,
"intermediate_size": 3072,
"num_attention_heads": 12,
"num_hidden_layers": 12,
"projection_dim": 768
},
"torch_dtype": "float32",
"transformers_version": null,
"vision_config": {
"_name_or_path": "",
"add_cross_attention": false,
"architectures": null,
"attention_dropout": 0.0,
"bad_words_ids": null,
"bos_token_id": null,
"chunk_size_feed_forward": 0,
"cross_attention_hidden_size": null,
"decoder_start_token_id": null,
"diversity_penalty": 0.0,
"do_sample": false,
"dropout": 0.0,
"early_stopping": false,
"encoder_no_repeat_ngram_size": 0,
"eos_token_id": null,
"finetuning_task": null,
"forced_bos_token_id": null,
"forced_eos_token_id": null,
"hidden_act": "quick_gelu",
"hidden_size": 1024,
"id2label": {
"0": "LABEL_0",
"1": "LABEL_1"
},
"image_size": 224,
"initializer_factor": 1.0,
"initializer_range": 0.02,
"intermediate_size": 4096,
"is_decoder": false,
"is_encoder_decoder": false,
"label2id": {
"LABEL_0": 0,
"LABEL_1": 1
},
"layer_norm_eps": 1e-05,
"length_penalty": 1.0,
"max_length": 20,
"min_length": 0,
"model_type": "clip_vision_model",
"no_repeat_ngram_size": 0,
"num_attention_heads": 16,
"num_beam_groups": 1,
"num_beams": 1,
"num_hidden_layers": 24,
"num_return_sequences": 1,
"output_attentions": false,
"output_hidden_states": false,
"output_scores": false,
"pad_token_id": null,
"patch_size": 14,
"prefix": null,
"problem_type": null,
"projection_dim" : 768,
"pruned_heads": {},
"remove_invalid_values": false,
"repetition_penalty": 1.0,
"return_dict": true,
"return_dict_in_generate": false,
"sep_token_id": null,
"task_specific_params": null,
"temperature": 1.0,
"tie_encoder_decoder": false,
"tie_word_embeddings": true,
"tokenizer_class": null,
"top_k": 50,
"top_p": 1.0,
"torch_dtype": null,
"torchscript": false,
"transformers_version": "4.16.0.dev0",
"use_bfloat16": false
},
"vision_config_dict": {
"hidden_size": 1024,
"intermediate_size": 4096,
"num_attention_heads": 16,
"num_hidden_layers": 24,
"patch_size": 14,
"projection_dim": 768
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,19 @@
{
"crop_size": 224,
"do_center_crop": true,
"do_normalize": true,
"do_resize": true,
"feature_extractor_type": "CLIPFeatureExtractor",
"image_mean": [
0.48145466,
0.4578275,
0.40821073
],
"image_std": [
0.26862954,
0.26130258,
0.27577711
],
"resample": 3,
"size": 224
}
@@ -0,0 +1 @@
{"bos_token": {"content": "<|startoftext|>", "single_word": false, "lstrip": false, "rstrip": false, "normalized": true}, "eos_token": {"content": "<|endoftext|>", "single_word": false, "lstrip": false, "rstrip": false, "normalized": true}, "unk_token": {"content": "<|endoftext|>", "single_word": false, "lstrip": false, "rstrip": false, "normalized": true}, "pad_token": "<|endoftext|>"}
File diff suppressed because it is too large Load Diff
+34
View File
@@ -0,0 +1,34 @@
{
"unk_token": {
"content": "<|endoftext|>",
"single_word": false,
"lstrip": false,
"rstrip": false,
"normalized": true,
"__type": "AddedToken"
},
"bos_token": {
"content": "<|startoftext|>",
"single_word": false,
"lstrip": false,
"rstrip": false,
"normalized": true,
"__type": "AddedToken"
},
"eos_token": {
"content": "<|endoftext|>",
"single_word": false,
"lstrip": false,
"rstrip": false,
"normalized": true,
"__type": "AddedToken"
},
"pad_token": "<|endoftext|>",
"add_prefix_space": false,
"errors": "replace",
"do_lower_case": true,
"name_or_path": "openai/clip-vit-base-patch32",
"model_max_length": 77,
"special_tokens_map_file": "./special_tokens_map.json",
"tokenizer_class": "CLIPTokenizer"
}
File diff suppressed because one or more lines are too long
+34
View File
@@ -0,0 +1,34 @@
{
"add_prefix_space": false,
"bos_token": {
"__type": "AddedToken",
"content": "<|startoftext|>",
"lstrip": false,
"normalized": true,
"rstrip": false,
"single_word": false
},
"do_lower_case": true,
"eos_token": {
"__type": "AddedToken",
"content": "<|endoftext|>",
"lstrip": false,
"normalized": true,
"rstrip": false,
"single_word": false
},
"errors": "replace",
"model_max_length": 77,
"name_or_path": "openai/clip-vit-large-patch14",
"pad_token": "<|endoftext|>",
"special_tokens_map_file": "./special_tokens_map.json",
"tokenizer_class": "CLIPTokenizer",
"unk_token": {
"__type": "AddedToken",
"content": "<|endoftext|>",
"lstrip": false,
"normalized": true,
"rstrip": false,
"single_word": false
}
}
+70
View File
@@ -0,0 +1,70 @@
model:
base_learning_rate: 1.0e-04
target: ldm.models.diffusion.ddpm.LatentDiffusion
params:
linear_start: 0.00085
linear_end: 0.0120
num_timesteps_cond: 1
log_every_t: 200
timesteps: 1000
first_stage_key: "jpg"
cond_stage_key: "txt"
image_size: 64
channels: 4
cond_stage_trainable: false # Note: different from the one we trained before
conditioning_key: crossattn
monitor: val/loss_simple_ema
scale_factor: 0.18215
use_ema: False
scheduler_config: # 10000 warmup steps
target: ldm.lr_scheduler.LambdaLinearScheduler
params:
warm_up_steps: [ 10000 ]
cycle_lengths: [ 10000000000000 ] # incredibly large number to prevent corner cases
f_start: [ 1.e-6 ]
f_max: [ 1. ]
f_min: [ 1. ]
unet_config:
target: ldm.modules.diffusionmodules.openaimodel.UNetModel
params:
image_size: 32 # unused
in_channels: 4
out_channels: 4
model_channels: 320
attention_resolutions: [ 4, 2, 1 ]
num_res_blocks: 2
channel_mult: [ 1, 2, 4, 4 ]
num_heads: 8
use_spatial_transformer: True
transformer_depth: 1
context_dim: 768
use_checkpoint: True
legacy: False
first_stage_config:
target: ldm.models.autoencoder.AutoencoderKL
params:
embed_dim: 4
monitor: val/rec_loss
ddconfig:
double_z: true
z_channels: 4
resolution: 256
in_channels: 3
out_ch: 3
ch: 128
ch_mult:
- 1
- 2
- 4
- 4
num_res_blocks: 2
attn_resolutions: []
dropout: 0.0
lossconfig:
target: torch.nn.Identity
cond_stage_config:
target: ldm.modules.encoders.modules.FrozenCLIPEmbedder
+419
View File
@@ -0,0 +1,419 @@
from pathlib import Path
from typing import Any, Optional, Union
import fire
import gradio as gr
import safetensors.torch
import torch
from diffusers import DPMSolverMultistepScheduler, StableDiffusionPipeline
from torchvision.utils import save_image
from model import ELLA, T5TextEmbedder
class ELLAProxyUNet(torch.nn.Module):
def __init__(self, ella, unet):
super().__init__()
# In order to still use the diffusers pipeline, including various workaround
self.ella = ella
self.unet = unet
self.config = unet.config
self.dtype = unet.dtype
self.device = unet.device
self.flexible_max_length_workaround = None
def forward(
self,
sample: torch.FloatTensor,
timestep: Union[torch.Tensor, float, int],
encoder_hidden_states: torch.Tensor,
class_labels: Optional[torch.Tensor] = None,
timestep_cond: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
cross_attention_kwargs: Optional[dict[str, Any]] = None,
added_cond_kwargs: Optional[dict[str, torch.Tensor]] = None,
down_block_additional_residuals: Optional[tuple[torch.Tensor]] = None,
mid_block_additional_residual: Optional[torch.Tensor] = None,
down_intrablock_additional_residuals: Optional[tuple[torch.Tensor]] = None,
encoder_attention_mask: Optional[torch.Tensor] = None,
return_dict: bool = True,
):
if self.flexible_max_length_workaround is not None:
time_aware_encoder_hidden_state_list = []
for i, max_length in enumerate(self.flexible_max_length_workaround):
time_aware_encoder_hidden_state_list.append(
self.ella(encoder_hidden_states[i : i + 1, :max_length], timestep)
)
# No matter how many tokens are text features, the ella output must be 64 tokens.
time_aware_encoder_hidden_states = torch.cat(
time_aware_encoder_hidden_state_list, dim=0
)
else:
time_aware_encoder_hidden_states = self.ella(
encoder_hidden_states, timestep
)
return self.unet(
sample=sample,
timestep=timestep,
encoder_hidden_states=time_aware_encoder_hidden_states,
class_labels=class_labels,
timestep_cond=timestep_cond,
attention_mask=attention_mask,
cross_attention_kwargs=cross_attention_kwargs,
added_cond_kwargs=added_cond_kwargs,
down_block_additional_residuals=down_block_additional_residuals,
mid_block_additional_residual=mid_block_additional_residual,
down_intrablock_additional_residuals=down_intrablock_additional_residuals,
encoder_attention_mask=encoder_attention_mask,
return_dict=return_dict,
)
def generate_image_with_flexible_max_length(
pipe, t5_encoder, prompt, fixed_negative=False, output_type="pt", **pipe_kwargs
):
device = pipe.device
dtype = pipe.dtype
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
prompt_embeds = t5_encoder(prompt, max_length=None).to(device, dtype)
negative_prompt_embeds = t5_encoder(
[""] * batch_size, max_length=128 if fixed_negative else None
).to(device, dtype)
# diffusers pipeline concatenate `prompt_embeds` too early...
# https://github.com/huggingface/diffusers/blob/b6d7e31d10df675d86c6fe7838044712c6dca4e9/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py#L913
pipe.unet.flexible_max_length_workaround = [
negative_prompt_embeds.size(1)
] * batch_size + [prompt_embeds.size(1)] * batch_size
max_length = max([prompt_embeds.size(1), negative_prompt_embeds.size(1)])
b, _, d = prompt_embeds.shape
prompt_embeds = torch.cat(
[
prompt_embeds,
torch.zeros(
(b, max_length - prompt_embeds.size(1), d), device=device, dtype=dtype
),
],
dim=1,
)
negative_prompt_embeds = torch.cat(
[
negative_prompt_embeds,
torch.zeros(
(b, max_length - negative_prompt_embeds.size(1), d),
device=device,
dtype=dtype,
),
],
dim=1,
)
images = pipe(
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
**pipe_kwargs,
output_type=output_type,
).images
pipe.unet.flexible_max_length_workaround = None
return images
def load_ella(filename, device, dtype):
ella = ELLA()
safetensors.torch.load_model(ella, filename, strict=True)
ella.to(device, dtype=dtype)
return ella
def load_ella_for_pipe(pipe, ella):
pipe.unet = ELLAProxyUNet(ella, pipe.unet)
def offload_ella_for_pipe(pipe):
pipe.unet = pipe.unet.unet
def generate_image_with_fixed_max_length(
pipe, t5_encoder, prompt, output_type="pt", **pipe_kwargs
):
prompt = [prompt] if isinstance(prompt, str) else prompt
prompt_embeds = t5_encoder(prompt, max_length=128).to(pipe.device, pipe.dtype)
negative_prompt_embeds = t5_encoder([""] * len(prompt), max_length=128).to(
pipe.device, pipe.dtype
)
return pipe(
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
**pipe_kwargs,
output_type=output_type,
).images
def build_demo(ella_path, sd_path="runwayml/stable-diffusion-v1-5"):
pipe = StableDiffusionPipeline.from_pretrained(
sd_path,
torch_dtype=torch.float16,
safety_checker=None,
feature_extractor=None,
requires_safety_checker=False,
)
pipe = pipe.to("cuda")
pipe.scheduler = DPMSolverMultistepScheduler.from_config(pipe.scheduler.config)
ella = load_ella(ella_path, pipe.device, pipe.dtype)
t5_encoder = T5TextEmbedder().to(pipe.device, dtype=torch.float16)
def generate_images(
prompt, guidance_scale, seed, num_inference_steps, size=512, _batch_size=2
):
print("#" * 50)
print(prompt)
load_ella_for_pipe(pipe, ella)
image_flexible = generate_image_with_flexible_max_length(
pipe,
t5_encoder,
[prompt] * _batch_size,
guidance_scale=guidance_scale,
num_inference_steps=num_inference_steps,
height=size,
width=size,
generator=[
torch.Generator(device="cuda").manual_seed(seed + i)
for i in range(_batch_size)
],
output_type="pil",
)
offload_ella_for_pipe(pipe)
image_ori = pipe(
[prompt] * _batch_size,
output_type="pil",
guidance_scale=guidance_scale,
num_inference_steps=num_inference_steps,
height=size,
width=size,
generator=[
torch.Generator(device="cuda").manual_seed(seed + i)
for i in range(_batch_size)
],
).images
return image_ori, image_flexible
with gr.Blocks() as app:
gr.Markdown(
"""
# ELLA-SD1.5 vs SD1.5
[ELLA Project](https://ella-diffusion.github.io/)
## Notes
** short prompt also works, but the result is much better after the caption is refined. **
### Caption Refining with In Context Learning(ICL)
caption refining instruction example:
```
Please generate the long prompt version of the short one according to the given examples. Long prompt version should consist of 3 to 5 sentences. Long prompt version must sepcify the color, shape, texture or spatial relation of the included objects. DO NOT generate sentences that describe any atmosphere!!!
Short: A calico cat with eyes closed is perched upon a Mercedes.
Long: a multicolored cat perched atop a shiny black car. the car is parked in front of a building with wooden walls and a green fence. the reflection of the car and the surrounding environment can be seen on the car's glossy surface.
Short: A boys sitting on a chair holding a video game remote.
Long: a young boy sitting on a chair, wearing a blue shirt and a baseball cap with the letter 'm'. he has a red medal around his neck and is holding a white game controller. behind him, there are two other individuals, one of whom is wearing a backpack. to the right of the boy, there's a blue trash bin with a sign that reads 'automatic party'.
Short: A man is on the bank of the water fishing.
Long: a serene waterscape where a person, dressed in a blue jacket and a red beanie, stands in shallow waters, fishing with a long rod. the calm waters are dotted with several sailboats anchored at a distance, and a mountain range can be seen in the background under a cloudy sky.
Short: A kitchen with a cluttered counter and wooden cabinets.
Long: a well-lit kitchen with wooden cabinets, a black and white checkered floor, and a refrigerator adorned with a floral decal on its side. the kitchen countertop holds various items, including a coffee maker, jars, and fruits.
Short: a racoon holding a shiny red apple over its head
```
using: https://huggingface.co/spaces/Qwen/Qwen-72B-Chat-Demo
got: a mischievous raccoon standing on its hind legs, holding a bright red apple aloft in its furry paws. the apple shines brightly against the backdrop of a dense forest, with leaves rustling in the gentle breeze. a few scattered rocks can be seen on the ground beneath the raccoon's feet, while a gnarled tree trunk stands nearby.
"""
)
with gr.Row():
input_caption = gr.Textbox(
value="A vivid red book with a smooth, matte cover lies next to a glossy yellow vase. The vase, with a slightly curved silhouette, stands on a dark wood table with a noticeable grain pattern. The book appears slightly worn at the edges, suggesting frequent use, while the vase holds a fresh array of multicolored wildflowers."
)
with gr.Column():
guidance_scale = gr.Slider(
minimum=1.0, maximum=16.0, value=10, label="guidance_scale"
)
seed = gr.Slider(
minimum=1000, maximum=2**20, value=1000, label="random seed"
)
num_inference_steps = gr.Slider(
minimum=15, maximum=100, value=25, label="num_inference_steps"
)
with gr.Row():
with gr.Column():
gr.Markdown(f"### ORIGINAL Stable Diffusion Model")
sd_output_image_gallery = gr.Gallery(columns=2, label="ORIGINAL SD")
with gr.Column():
gr.Markdown(f"### ELLA")
ella_output_image_gallery = gr.Gallery(columns=2, label="ELLA")
submit_button = gr.Button()
submit_button.click(
fn=generate_images,
inputs=[input_caption, guidance_scale, seed, num_inference_steps],
outputs=[sd_output_image_gallery, ella_output_image_gallery],
)
app.queue(concurrency_count=1, api_open=False)
app.launch(share=False)
def main(save_folder, ella_path):
save_folder = Path(save_folder)
save_folder.mkdir(exist_ok=True)
pipe = StableDiffusionPipeline.from_pretrained(
"runwayml/stable-diffusion-v1-5",
torch_dtype=torch.float16,
safety_checker=None,
feature_extractor=None,
requires_safety_checker=False,
)
pipe = pipe.to("cuda")
pipe.scheduler = DPMSolverMultistepScheduler.from_config(pipe.scheduler.config)
ella = load_ella(ella_path, pipe.device, pipe.dtype)
t5_encoder = T5TextEmbedder().to(pipe.device, dtype=torch.float16)
# prompt from ViLG-300, PartiPrompts
# short prompt also works, but the result is much better after the caption is refined.
# caption refining instruction example:
# ```
# Please generate the long prompt version of the short one according to the given examples. Long prompt version should consist of 3 to 5 sentences. Long prompt version must sepcify the color, shape, texture or spatial relation of the included objects. DO NOT generate sentences that describe any atmosphere!!!
#
# Short: A calico cat with eyes closed is perched upon a Mercedes.
# Long: a multicolored cat perched atop a shiny black car. the car is parked in front of a building with wooden walls and a green fence. the reflection of the car and the surrounding environment can be seen on the car's glossy surface.
#
# Short: A boys sitting on a chair holding a video game remote.
# Long: a young boy sitting on a chair, wearing a blue shirt and a baseball cap with the letter 'm'. he has a red medal around his neck and is holding a white game controller. behind him, there are two other individuals, one of whom is wearing a backpack. to the right of the boy, there's a blue trash bin with a sign that reads 'automatic party'.
#
# Short: A man is on the bank of the water fishing.
# Long: a serene waterscape where a person, dressed in a blue jacket and a red beanie, stands in shallow waters, fishing with a long rod. the calm waters are dotted with several sailboats anchored at a distance, and a mountain range can be seen in the background under a cloudy sky.
#
# Short: A kitchen with a cluttered counter and wooden cabinets.
# Long: a well-lit kitchen with wooden cabinets, a black and white checkered floor, and a refrigerator adorned with a floral decal on its side. the kitchen countertop holds various items, including a coffee maker, jars, and fruits.
#
# Short: a racoon holding a shiny red apple over its head
# ```
#
# using: https://huggingface.co/spaces/Qwen/Qwen-72B-Chat-Demo
# got: a mischievous raccoon standing on its hind legs, holding a bright red apple aloft in its furry paws. the apple shines brightly against the backdrop of a dense forest, with leaves rustling in the gentle breeze. a few scattered rocks can be seen on the ground beneath the raccoon's feet, while a gnarled tree trunk stands nearby.
prompt_name_examples1 = [
("crocodile_sweater", "Crocodile in a sweater"),
(
"crocodile_sweater-gpt4_refined_caption",
"a large, textured green crocodile lying comfortably on a patch of grass with a cute, knitted orange sweater enveloping its scaly body. Around its neck, the sweater features a whimsical pattern of blue and yellow stripes. In the background, a smooth, grey rock partially obscures the view of a small pond with lily pads floating on the surface.",
),
("red_book-yellow_vase", "A red book and a yellow vase."),
(
"red_book-yellow_vase-gpt4_refined_caption",
"A vivid red book with a smooth, matte cover lies next to a glossy yellow vase. The vase, with a slightly curved silhouette, stands on a dark wood table with a noticeable grain pattern. The book appears slightly worn at the edges, suggesting frequent use, while the vase holds a fresh array of multicolored wildflowers.",
),
("racoon_apple", "a racoon holding a shiny red apple over its head"),
(
"racoon_apple_Qwen-72B-Chat-refined",
"a mischievous raccoon standing on its hind legs, holding a bright red apple aloft in its furry paws. the apple shines brightly against the backdrop of a dense forest, with leaves rustling in the gentle breeze. a few scattered rocks can be seen on the ground beneath the raccoon's feet, while a gnarled tree trunk stands nearby.",
),
]
# hard example prompt.
prompt_name_examples2 = [
(
"falcon_chinese",
"a chinese man wearing a white shirt and a checkered headscarf, holds a large falcon near his shoulder. the falcon has dark feathers with a distinctive beak. the background consists of a clear sky and a fence, suggesting an outdoor setting, possibly a desert or arid region",
),
(
"wombat",
"A close-up photo of a wombat wearing a red backpack and raising both arms in the air. Mount Rushmore is in the background",
),
(
"bakkot_AstralCodexTen_2",
"An oil painting of a man in a factory looking at a cat wearing a top hat",
),
]
for name, prompt in prompt_name_examples1 + prompt_name_examples2:
print("#" * 80)
print(f'{name}: "{prompt}"')
_batch_size = 1
size = 512
seed = 1001
prompt = [prompt] * _batch_size
load_ella_for_pipe(pipe, ella)
image_flexible = generate_image_with_flexible_max_length(
pipe,
t5_encoder,
prompt,
guidance_scale=12,
num_inference_steps=50,
height=size,
width=size,
generator=[
torch.Generator(device="cuda").manual_seed(seed + i)
for i in range(_batch_size)
],
)
image_fixed = generate_image_with_fixed_max_length(
pipe,
t5_encoder,
prompt,
guidance_scale=12,
num_inference_steps=50,
height=size,
width=size,
generator=[
torch.Generator(device="cuda").manual_seed(seed + i)
for i in range(_batch_size)
],
)
offload_ella_for_pipe(pipe)
image_ori = pipe(
prompt,
output_type="pt",
guidance_scale=12,
num_inference_steps=50,
height=size,
width=size,
generator=[
torch.Generator(device="cuda").manual_seed(seed + i)
for i in range(_batch_size)
],
).images
print(f'save image at {save_folder / f"{name}.png"}')
print(
"original SD1.5\t|\tELLA-SD1.5(fixed token length)\t|\tELLA-SD1.5(flexible token length)"
)
save_image(
torch.cat([image_ori, image_fixed, image_flexible], dim=0),
save_folder / f"{name}.png",
nrow=3,
)
if __name__ == "__main__":
fire.Fire(dict(test=main, demo=build_demo))
+218
View File
@@ -0,0 +1,218 @@
from collections import OrderedDict
from typing import Optional
import torch
import torch.nn as nn
from diffusers.models.embeddings import TimestepEmbedding, Timesteps
from transformers import T5EncoderModel, T5Tokenizer
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)
x = self.norm(x) * (1 + scale) + shift
return x
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: 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),
)
latents = latents + self.mlp(self.ln_ff(latents, timestep_embedding))
return latents
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)
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)
)
def forward(self, x: torch.Tensor, timestep_embedding: torch.Tensor = None):
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__()
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
):
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)
attention_mask = attention_mask.to(self.model.device)
outputs = self.model(text_input_ids, attention_mask=attention_mask)
embeddings = outputs.last_hidden_state
return embeddings
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__()
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,
)
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, text_encode_features, timesteps):
device = text_encode_features.device
dtype = text_encode_features.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(text_encode_features), -1, -1)
time_embedding = self.time_embedding(ori_time_feature)
encoder_hidden_states = self.connector(
text_encode_features, timestep_embedding=time_embedding
)
return encoder_hidden_states
+397
View File
@@ -0,0 +1,397 @@
import os
from typing import Optional
from typing import Any, Optional, Union
from contextlib import nullcontext
import safetensors.torch
import torch
from diffusers import DPMSolverMultistepScheduler, StableDiffusionPipeline, AutoencoderKL, UNet2DConditionModel, DDIMScheduler, LCMScheduler, DDPMScheduler, DEISMultistepScheduler, PNDMScheduler
from omegaconf import OmegaConf
from .model import ELLA, T5TextEmbedder
from transformers import CLIPTokenizer
import comfy.model_management as mm
import comfy.utils
script_directory = os.path.dirname(os.path.abspath(__file__))
class ELLAProxyUNet(torch.nn.Module):
def __init__(self, ella, unet):
super().__init__()
# In order to still use the diffusers pipeline, including various workaround
self.ella = ella
self.unet = unet
self.config = unet.config
self.dtype = unet.dtype
self.device = unet.device
self.flexible_max_length_workaround = None
def forward(
self,
sample: torch.FloatTensor,
timestep: Union[torch.Tensor, float, int],
encoder_hidden_states: torch.Tensor,
class_labels: Optional[torch.Tensor] = None,
timestep_cond: Optional[torch.Tensor] = None,
attention_mask: Optional[torch.Tensor] = None,
cross_attention_kwargs: Optional[dict[str, Any]] = None,
added_cond_kwargs: Optional[dict[str, torch.Tensor]] = None,
down_block_additional_residuals: Optional[tuple[torch.Tensor]] = None,
mid_block_additional_residual: Optional[torch.Tensor] = None,
down_intrablock_additional_residuals: Optional[tuple[torch.Tensor]] = None,
encoder_attention_mask: Optional[torch.Tensor] = None,
return_dict: bool = True,
):
if self.flexible_max_length_workaround is not None:
time_aware_encoder_hidden_state_list = []
for i, max_length in enumerate(self.flexible_max_length_workaround):
time_aware_encoder_hidden_state_list.append(
self.ella(encoder_hidden_states[i : i + 1, :max_length], timestep)
)
# No matter how many tokens are text features, the ella output must be 64 tokens.
time_aware_encoder_hidden_states = torch.cat(
time_aware_encoder_hidden_state_list, dim=0
)
else:
time_aware_encoder_hidden_states = self.ella(
encoder_hidden_states, timestep
)
return self.unet(
sample=sample,
timestep=timestep,
encoder_hidden_states=time_aware_encoder_hidden_states,
class_labels=class_labels,
timestep_cond=timestep_cond,
attention_mask=attention_mask,
cross_attention_kwargs=cross_attention_kwargs,
added_cond_kwargs=added_cond_kwargs,
down_block_additional_residuals=down_block_additional_residuals,
mid_block_additional_residual=mid_block_additional_residual,
down_intrablock_additional_residuals=down_intrablock_additional_residuals,
encoder_attention_mask=encoder_attention_mask,
return_dict=return_dict,
)
def generate_image_with_flexible_max_length(
pipe, t5_encoder, prompt, fixed_negative=False, output_type="pt", **pipe_kwargs
):
device = pipe.device
dtype = pipe.dtype
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
prompt_embeds = t5_encoder(prompt, max_length=None).to(device, dtype)
negative_prompt_embeds = t5_encoder(
[""] * batch_size, max_length=128 if fixed_negative else None
).to(device, dtype)
# diffusers pipeline concatenate `prompt_embeds` too early...
# https://github.com/huggingface/diffusers/blob/b6d7e31d10df675d86c6fe7838044712c6dca4e9/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py#L913
pipe.unet.flexible_max_length_workaround = [
negative_prompt_embeds.size(1)
] * batch_size + [prompt_embeds.size(1)] * batch_size
max_length = max([prompt_embeds.size(1), negative_prompt_embeds.size(1)])
b, _, d = prompt_embeds.shape
prompt_embeds = torch.cat(
[
prompt_embeds,
torch.zeros(
(b, max_length - prompt_embeds.size(1), d), device=device, dtype=dtype
),
],
dim=1,
)
negative_prompt_embeds = torch.cat(
[
negative_prompt_embeds,
torch.zeros(
(b, max_length - negative_prompt_embeds.size(1), d),
device=device,
dtype=dtype,
),
],
dim=1,
)
images = pipe(
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
**pipe_kwargs,
output_type=output_type,
).images
pipe.unet.flexible_max_length_workaround = None
return images
def load_ella(filename, device, dtype):
ella = ELLA()
safetensors.torch.load_model(ella, filename, strict=True)
ella.to(device, dtype=dtype)
return ella
def load_ella_for_pipe(pipe, ella):
pipe.unet = ELLAProxyUNet(ella, pipe.unet)
def offload_ella_for_pipe(pipe):
pipe.unet = pipe.unet.unet
def generate_image_with_fixed_max_length(
pipe, t5_encoder, prompt, output_type="pt", **pipe_kwargs
):
prompt = [prompt] if isinstance(prompt, str) else prompt
prompt_embeds = t5_encoder(prompt, max_length=128).to(pipe.device, pipe.dtype)
negative_prompt_embeds = t5_encoder([""] * len(prompt), max_length=128).to(
pipe.device, pipe.dtype
)
return pipe(
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
**pipe_kwargs,
output_type=output_type,
).images
class ella_model_loader:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": ("MODEL",),
"clip": ("CLIP",),
"vae": ("VAE",),
},
}
RETURN_TYPES = ("ELLAMODEL",)
RETURN_NAMES = ("ella_model",)
FUNCTION = "loadmodel"
CATEGORY = "ellaWrapper"
def loadmodel(self, model, clip, vae):
mm.soft_empty_cache()
dtype = mm.unet_dtype()
device = mm.get_torch_device()
custom_config = {
'model': model,
'vae': vae,
}
if not hasattr(self, 'model') or self.model == None or custom_config != self.current_config:
pbar = comfy.utils.ProgressBar(5)
self.current_config = custom_config
# setup pretrained models
original_config = OmegaConf.load(os.path.join(script_directory, f"configs/v1-inference.yaml"))
print("loading ELLA")
checkpoint_path = os.path.join(script_directory, 'checkpoints')
ella_path = os.path.join(checkpoint_path, 'ella-sd1.5-tsc-t5xl.safetensors')
if not os.path.exists(ella_path):
from huggingface_hub import snapshot_download
snapshot_download(repo_id="QQGYLab/ELLA", local_dir=checkpoint_path, local_dir_use_symlinks=False)
from diffusers.loaders.single_file_utils import (convert_ldm_vae_checkpoint, convert_ldm_unet_checkpoint, create_text_encoder_from_ldm_clip_checkpoint, create_vae_diffusers_config, create_unet_diffusers_config)
ella = ELLA()
safetensors.torch.load_model(ella, ella_path, strict=True)
clip_sd = None
load_models = [model]
load_models.append(clip.load_model())
clip_sd = clip.get_sd()
comfy.model_management.load_models_gpu(load_models)
sd = model.model.state_dict_for_saving(clip_sd, vae.get_sd(), None)
# 1. vae
converted_vae_config = create_vae_diffusers_config(original_config, image_size=512)
converted_vae = convert_ldm_vae_checkpoint(sd, converted_vae_config)
vae = AutoencoderKL(**converted_vae_config)
vae.load_state_dict(converted_vae, strict=False)
pbar.update(1)
# 2. unet
converted_unet_config = create_unet_diffusers_config(original_config, image_size=512)
converted_unet = convert_ldm_unet_checkpoint(sd, converted_unet_config)
unet = UNet2DConditionModel(**converted_unet_config)
unet.load_state_dict(converted_unet, strict=False)
pbar.update(1)
# 3. text_model
print("loading text model")
text_encoder = create_text_encoder_from_ldm_clip_checkpoint("openai/clip-vit-large-patch14",sd)
scheduler_config = {
'num_train_timesteps': 1000,
'beta_start': 0.00085,
'beta_end': 0.012,
'beta_schedule': "linear",
'steps_offset': 1
}
scheduler=DPMSolverMultistepScheduler(**scheduler_config)
pbar.update(1)
del sd
print("loading ELLA")
ella_path = os.path.join(script_directory, 'checkpoints', 'ella-sd1.5-tsc-t5xl.safetensors')
ella = ELLA()
safetensors.torch.load_model(ella, ella_path, strict=True)
ella.to(device, dtype=dtype)
unet = unet.to(device)
ella_unet = ELLAProxyUNet(ella, unet)
pbar.update(1)
print("loading tokenizer")
tokenizer_path = os.path.join(script_directory, "configs/tokenizer")
tokenizer = CLIPTokenizer.from_pretrained(tokenizer_path)
print("creating pipeline")
pipe = StableDiffusionPipeline(
unet=unet,
vae=vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
scheduler=scheduler,
safety_checker=None,
feature_extractor=None,
requires_safety_checker=False,
image_encoder=None
)
print("pipeline created")
pbar.update(1)
pipe.unet = ella_unet
t5_encoder = T5TextEmbedder().to(pipe.device, dtype=torch.float16)
ella_model = {
'pipe': pipe,
'ella': ella,
't5_encoder': t5_encoder
}
return (ella_model,)
class ella_sampler:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"ella_model": ("ELLAMODEL",),
"prompt": ("STRING", {"multiline": True, "default": "A vivid red book with a smooth, matte cover lies next to a glossy yellow vase. The vase, with a slightly curved silhouette, stands on a dark wood table with a noticeable grain pattern. The book appears slightly worn at the edges, suggesting frequent use, while the vase holds a fresh array of multicolored wildflowers.",}),
"width": ("INT", {"default": 512, "min": 64, "max": 2048, "step": 64}),
"height": ("INT", {"default": 512, "min": 64, "max": 2048, "step": 64}),
"batch_size": ("INT", {"default": 1, "min": 1, "max": 256, "step": 1}),
"steps": ("INT", {"default": 25, "min": 1, "max": 200, "step": 1}),
"guidance_scale": ("FLOAT", {"default": 10.0, "min": 0.0, "max": 20.0, "step": 0.01}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"scheduler": (
[
'DDIMScheduler',
'DDPMScheduler',
'LCMScheduler',
'PNDMScheduler',
'DEISMultistepScheduler'
], {
"default": 'DDIMScheduler'
}),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("images",)
FUNCTION = "process"
CATEGORY = "champWrapper"
def process(self, prompt, batch_size, width, height, steps, guidance_scale, seed, ella_model, scheduler):
device = mm.get_torch_device()
mm.unload_all_models()
mm.soft_empty_cache()
dtype = mm.unet_dtype()
t5_encoder=ella_model['t5_encoder']
pipe=ella_model['pipe']
pipe.to(device, dtype=dtype)
scheduler_config = {
'num_train_timesteps': 1000,
'beta_start': 0.00085,
'beta_end': 0.012,
'beta_schedule': "linear",
'steps_offset': 1
}
if scheduler == 'DDIMScheduler':
noise_scheduler = DDIMScheduler(**scheduler_config)
elif scheduler == 'DDPMScheduler':
noise_scheduler = DDPMScheduler(**scheduler_config)
elif scheduler == 'LCMScheduler':
noise_scheduler = LCMScheduler(**scheduler_config)
elif scheduler == 'PNDMScheduler':
noise_scheduler = PNDMScheduler(**scheduler_config)
elif scheduler == 'DEISMultistepScheduler':
noise_scheduler = DEISMultistepScheduler(**scheduler_config)
pipe.scheduler = noise_scheduler
autocast_condition = (dtype != torch.float32) and not mm.is_device_mps(device)
with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
fixed_negative = False
prompt_embeds = t5_encoder(prompt, max_length=None).to(device, dtype)
negative_prompt_embeds = t5_encoder(
[""] * batch_size, max_length=128 if fixed_negative else None
).to(device, dtype)
pipe.unet.flexible_max_length_workaround = [
negative_prompt_embeds.size(1)
] * batch_size + [prompt_embeds.size(1)] * batch_size
max_length = max([prompt_embeds.size(1), negative_prompt_embeds.size(1)])
b, _, d = prompt_embeds.shape
prompt_embeds = torch.cat(
[
prompt_embeds,
torch.zeros(
(b, max_length - prompt_embeds.size(1), d), device=device, dtype=dtype
),
],
dim=1,
)
negative_prompt_embeds = torch.cat(
[
negative_prompt_embeds,
torch.zeros(
(b, max_length - negative_prompt_embeds.size(1), d),
device=device,
dtype=dtype,
),
],
dim=1,
)
images = pipe(
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
guidance_scale=guidance_scale,
num_inference_steps=steps,
height=height,
width=width,
generator=[
torch.Generator(device=device).manual_seed(seed + i)
for i in range(batch_size)
],
output_type="np.array",
).images
print(images.shape)
tensor = torch.from_numpy(images).cpu().float()
return (tensor,)
NODE_CLASS_MAPPINGS = {
"ella_model_loader": ella_model_loader,
"ella_sampler": ella_sampler,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ella_model_loader": "ELLA Model Loader",
"ella_sampler": "ELLA Sampler",
}
+2
View File
@@ -0,0 +1,2 @@
diffusers>=0.26.0
omegaconfg