Initial commit
This commit is contained in:
@@ -0,0 +1,6 @@
|
||||
__pycache__/
|
||||
checkpoints/
|
||||
*.py[cod]
|
||||
*$py.class
|
||||
*.egg-info
|
||||
.pytest_cache
|
||||
@@ -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.
|
||||
@@ -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✦</a>
|
||||
</span>
|
||||
<p>
|
||||
* Equal contributions, ✦ 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}
|
||||
}
|
||||
```
|
||||
@@ -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 |
Executable
+25
@@ -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
|
||||
}
|
||||
@@ -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
@@ -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
Executable
+34
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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))
|
||||
@@ -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
|
||||
@@ -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",
|
||||
}
|
||||
@@ -0,0 +1,2 @@
|
||||
diffusers>=0.26.0
|
||||
omegaconfg
|
||||
Reference in New Issue
Block a user