Instructions to use BiliSakura/Looped-DiT-diffusers with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use BiliSakura/Looped-DiT-diffusers with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("BiliSakura/Looped-DiT-diffusers", dtype=torch.bfloat16, device_map="cuda") prompt = "a red cube on top of a blue sphere" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- Draw Things
- DiffusionBee
Add files using upload-large-folder tool
Browse files- .gitattributes +37 -35
- Looped-DiT-B-16/demo.png +3 -0
- Looped-DiT-B-16/model_index.json +23 -0
- Looped-DiT-B-16/pipeline.py +763 -0
- Looped-DiT-B-16/scheduler/scheduler_config.json +18 -0
- Looped-DiT-B-16/text_encoder/README.md +276 -0
- Looped-DiT-B-16/text_encoder/config.json +28 -0
- Looped-DiT-B-16/text_encoder/generation_config.json +7 -0
- Looped-DiT-B-16/text_encoder/model.safetensors +3 -0
- Looped-DiT-B-16/text_encoder/special_tokens_map.json +107 -0
- Looped-DiT-B-16/text_encoder/spiece.model +3 -0
- Looped-DiT-B-16/text_encoder/tokenizer.json +0 -0
- Looped-DiT-B-16/text_encoder/tokenizer_config.json +113 -0
- Looped-DiT-B-16/tokenizer/special_tokens_map.json +107 -0
- Looped-DiT-B-16/tokenizer/spiece.model +3 -0
- Looped-DiT-B-16/tokenizer/tokenizer.json +0 -0
- Looped-DiT-B-16/tokenizer/tokenizer_config.json +113 -0
- Looped-DiT-B-16/transformer/config.json +23 -0
- Looped-DiT-B-16/transformer/diffusion_pytorch_model.safetensors +3 -0
- Looped-DiT-B-16/transformer/transformer_looped_dit.py +417 -0
- Looped-DiT-B-32/demo.png +3 -0
- Looped-DiT-B-32/model_index.json +23 -0
- Looped-DiT-B-32/pipeline.py +763 -0
- Looped-DiT-B-32/scheduler/scheduler_config.json +18 -0
- Looped-DiT-B-32/text_encoder/README.md +276 -0
- Looped-DiT-B-32/text_encoder/config.json +28 -0
- Looped-DiT-B-32/text_encoder/generation_config.json +7 -0
- Looped-DiT-B-32/text_encoder/model.safetensors +3 -0
- Looped-DiT-B-32/text_encoder/special_tokens_map.json +107 -0
- Looped-DiT-B-32/text_encoder/spiece.model +3 -0
- Looped-DiT-B-32/text_encoder/tokenizer.json +0 -0
- Looped-DiT-B-32/text_encoder/tokenizer_config.json +113 -0
- Looped-DiT-B-32/tokenizer/special_tokens_map.json +107 -0
- Looped-DiT-B-32/tokenizer/spiece.model +3 -0
- Looped-DiT-B-32/tokenizer/tokenizer.json +0 -0
- Looped-DiT-B-32/tokenizer/tokenizer_config.json +113 -0
- Looped-DiT-B-32/transformer/config.json +23 -0
- Looped-DiT-B-32/transformer/diffusion_pytorch_model.safetensors +3 -0
- Looped-DiT-B-32/transformer/transformer_looped_dit.py +417 -0
- README.md +151 -0
.gitattributes
CHANGED
|
@@ -1,35 +1,37 @@
|
|
| 1 |
-
*.7z filter=lfs diff=lfs merge=lfs -text
|
| 2 |
-
*.arrow filter=lfs diff=lfs merge=lfs -text
|
| 3 |
-
*.bin filter=lfs diff=lfs merge=lfs -text
|
| 4 |
-
*.bz2 filter=lfs diff=lfs merge=lfs -text
|
| 5 |
-
*.ckpt filter=lfs diff=lfs merge=lfs -text
|
| 6 |
-
*.ftz filter=lfs diff=lfs merge=lfs -text
|
| 7 |
-
*.gz filter=lfs diff=lfs merge=lfs -text
|
| 8 |
-
*.h5 filter=lfs diff=lfs merge=lfs -text
|
| 9 |
-
*.joblib filter=lfs diff=lfs merge=lfs -text
|
| 10 |
-
*.lfs.* filter=lfs diff=lfs merge=lfs -text
|
| 11 |
-
*.mlmodel filter=lfs diff=lfs merge=lfs -text
|
| 12 |
-
*.model filter=lfs diff=lfs merge=lfs -text
|
| 13 |
-
*.msgpack filter=lfs diff=lfs merge=lfs -text
|
| 14 |
-
*.npy filter=lfs diff=lfs merge=lfs -text
|
| 15 |
-
*.npz filter=lfs diff=lfs merge=lfs -text
|
| 16 |
-
*.onnx filter=lfs diff=lfs merge=lfs -text
|
| 17 |
-
*.ot filter=lfs diff=lfs merge=lfs -text
|
| 18 |
-
*.parquet filter=lfs diff=lfs merge=lfs -text
|
| 19 |
-
*.pb filter=lfs diff=lfs merge=lfs -text
|
| 20 |
-
*.pickle filter=lfs diff=lfs merge=lfs -text
|
| 21 |
-
*.pkl filter=lfs diff=lfs merge=lfs -text
|
| 22 |
-
*.pt filter=lfs diff=lfs merge=lfs -text
|
| 23 |
-
*.pth filter=lfs diff=lfs merge=lfs -text
|
| 24 |
-
*.rar filter=lfs diff=lfs merge=lfs -text
|
| 25 |
-
*.safetensors filter=lfs diff=lfs merge=lfs -text
|
| 26 |
-
saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
| 27 |
-
*.tar.* filter=lfs diff=lfs merge=lfs -text
|
| 28 |
-
*.tar filter=lfs diff=lfs merge=lfs -text
|
| 29 |
-
*.tflite filter=lfs diff=lfs merge=lfs -text
|
| 30 |
-
*.tgz filter=lfs diff=lfs merge=lfs -text
|
| 31 |
-
*.wasm filter=lfs diff=lfs merge=lfs -text
|
| 32 |
-
*.xz filter=lfs diff=lfs merge=lfs -text
|
| 33 |
-
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
-
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
-
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
|
|
|
| 1 |
+
*.7z filter=lfs diff=lfs merge=lfs -text
|
| 2 |
+
*.arrow filter=lfs diff=lfs merge=lfs -text
|
| 3 |
+
*.bin filter=lfs diff=lfs merge=lfs -text
|
| 4 |
+
*.bz2 filter=lfs diff=lfs merge=lfs -text
|
| 5 |
+
*.ckpt filter=lfs diff=lfs merge=lfs -text
|
| 6 |
+
*.ftz filter=lfs diff=lfs merge=lfs -text
|
| 7 |
+
*.gz filter=lfs diff=lfs merge=lfs -text
|
| 8 |
+
*.h5 filter=lfs diff=lfs merge=lfs -text
|
| 9 |
+
*.joblib filter=lfs diff=lfs merge=lfs -text
|
| 10 |
+
*.lfs.* filter=lfs diff=lfs merge=lfs -text
|
| 11 |
+
*.mlmodel filter=lfs diff=lfs merge=lfs -text
|
| 12 |
+
*.model filter=lfs diff=lfs merge=lfs -text
|
| 13 |
+
*.msgpack filter=lfs diff=lfs merge=lfs -text
|
| 14 |
+
*.npy filter=lfs diff=lfs merge=lfs -text
|
| 15 |
+
*.npz filter=lfs diff=lfs merge=lfs -text
|
| 16 |
+
*.onnx filter=lfs diff=lfs merge=lfs -text
|
| 17 |
+
*.ot filter=lfs diff=lfs merge=lfs -text
|
| 18 |
+
*.parquet filter=lfs diff=lfs merge=lfs -text
|
| 19 |
+
*.pb filter=lfs diff=lfs merge=lfs -text
|
| 20 |
+
*.pickle filter=lfs diff=lfs merge=lfs -text
|
| 21 |
+
*.pkl filter=lfs diff=lfs merge=lfs -text
|
| 22 |
+
*.pt filter=lfs diff=lfs merge=lfs -text
|
| 23 |
+
*.pth filter=lfs diff=lfs merge=lfs -text
|
| 24 |
+
*.rar filter=lfs diff=lfs merge=lfs -text
|
| 25 |
+
*.safetensors filter=lfs diff=lfs merge=lfs -text
|
| 26 |
+
saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
| 27 |
+
*.tar.* filter=lfs diff=lfs merge=lfs -text
|
| 28 |
+
*.tar filter=lfs diff=lfs merge=lfs -text
|
| 29 |
+
*.tflite filter=lfs diff=lfs merge=lfs -text
|
| 30 |
+
*.tgz filter=lfs diff=lfs merge=lfs -text
|
| 31 |
+
*.wasm filter=lfs diff=lfs merge=lfs -text
|
| 32 |
+
*.xz filter=lfs diff=lfs merge=lfs -text
|
| 33 |
+
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
+
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
+
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
| 36 |
+
Looped-DiT-B-16/demo.png filter=lfs diff=lfs merge=lfs -text
|
| 37 |
+
Looped-DiT-B-32/demo.png filter=lfs diff=lfs merge=lfs -text
|
Looped-DiT-B-16/demo.png
ADDED
|
Git LFS Details
|
Looped-DiT-B-16/model_index.json
ADDED
|
@@ -0,0 +1,23 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"_class_name": [
|
| 3 |
+
"pipeline",
|
| 4 |
+
"LoopedDiTPipeline"
|
| 5 |
+
],
|
| 6 |
+
"_diffusers_version": "0.39.0",
|
| 7 |
+
"scheduler": [
|
| 8 |
+
"diffusers",
|
| 9 |
+
"FlowMatchEulerDiscreteScheduler"
|
| 10 |
+
],
|
| 11 |
+
"text_encoder": [
|
| 12 |
+
"transformers",
|
| 13 |
+
"T5EncoderModel"
|
| 14 |
+
],
|
| 15 |
+
"tokenizer": [
|
| 16 |
+
"transformers",
|
| 17 |
+
"T5Tokenizer"
|
| 18 |
+
],
|
| 19 |
+
"transformer": [
|
| 20 |
+
"transformer_looped_dit",
|
| 21 |
+
"LoopedDiTTransformer2DModel"
|
| 22 |
+
]
|
| 23 |
+
}
|
Looped-DiT-B-16/pipeline.py
ADDED
|
@@ -0,0 +1,763 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2025 The HuggingFace Team. All rights reserved.
|
| 2 |
+
#
|
| 3 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 4 |
+
# you may not use this file except in compliance with the License.
|
| 5 |
+
# You may obtain a copy of the License at
|
| 6 |
+
#
|
| 7 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 8 |
+
#
|
| 9 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 10 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 11 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 12 |
+
# See the License for the specific language governing permissions and
|
| 13 |
+
# limitations under the License.
|
| 14 |
+
|
| 15 |
+
import inspect
|
| 16 |
+
from typing import Any, Callable
|
| 17 |
+
|
| 18 |
+
import torch
|
| 19 |
+
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
|
| 20 |
+
from diffusers.models.modeling_utils import ModelMixin
|
| 21 |
+
from diffusers.pipelines.pipeline_utils import DiffusionPipeline, ImagePipelineOutput
|
| 22 |
+
from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion import retrieve_timesteps
|
| 23 |
+
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler, KarrasDiffusionSchedulers
|
| 24 |
+
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
| 25 |
+
from diffusers.utils import deprecate, is_torch_xla_available, logging, replace_example_docstring
|
| 26 |
+
from diffusers.utils.torch_utils import randn_tensor
|
| 27 |
+
from PIL import Image
|
| 28 |
+
from transformers import AutoTokenizer, T5EncoderModel
|
| 29 |
+
|
| 30 |
+
if is_torch_xla_available():
|
| 31 |
+
import torch_xla.core.xla_model as xm
|
| 32 |
+
|
| 33 |
+
XLA_AVAILABLE = True
|
| 34 |
+
else:
|
| 35 |
+
XLA_AVAILABLE = False
|
| 36 |
+
|
| 37 |
+
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
| 38 |
+
|
| 39 |
+
# Training clamps the flow-matching denominator so the loss stays finite at t -> 1.
|
| 40 |
+
VELOCITY_DENOM_MIN = 0.05
|
| 41 |
+
|
| 42 |
+
EXAMPLE_DOC_STRING = """
|
| 43 |
+
Examples:
|
| 44 |
+
```py
|
| 45 |
+
>>> from pathlib import Path
|
| 46 |
+
>>> import torch
|
| 47 |
+
>>> from diffusers import DiffusionPipeline
|
| 48 |
+
|
| 49 |
+
>>> model_dir = Path("checkpoints/looped-dit-b16").resolve()
|
| 50 |
+
>>> pipe = DiffusionPipeline.from_pretrained(
|
| 51 |
+
... str(model_dir),
|
| 52 |
+
... local_files_only=True,
|
| 53 |
+
... custom_pipeline=str(model_dir / "pipeline.py"),
|
| 54 |
+
... trust_remote_code=True,
|
| 55 |
+
... torch_dtype=torch.bfloat16,
|
| 56 |
+
... ).to("cuda")
|
| 57 |
+
|
| 58 |
+
>>> image = pipe(
|
| 59 |
+
... "a red cube on top of a blue sphere",
|
| 60 |
+
... num_inference_steps=100,
|
| 61 |
+
... guidance_scale=6.0,
|
| 62 |
+
... num_loops=4,
|
| 63 |
+
... generator=torch.Generator(device="cuda").manual_seed(0),
|
| 64 |
+
... ).images[0]
|
| 65 |
+
>>> image.save("sample.png")
|
| 66 |
+
|
| 67 |
+
>>> # Hugging Face Hub style model id: UserID/RepoID
|
| 68 |
+
>>> # RepoID is usually like "modelname-diffusers"
|
| 69 |
+
>>> # Example: "your-user/Looped-DiT-diffusers"
|
| 70 |
+
```
|
| 71 |
+
"""
|
| 72 |
+
|
| 73 |
+
|
| 74 |
+
def paper_euler_sigmas(num_inference_steps: int) -> list[float]:
|
| 75 |
+
r"""
|
| 76 |
+
Sigma grid of the training Euler sampler.
|
| 77 |
+
|
| 78 |
+
Training integrates flow time `t` from 0 (noise) to 1 (data) with
|
| 79 |
+
`torch.linspace(0, 1, steps + 1)`. Flow-match schedulers step in sigma
|
| 80 |
+
`1 - t` and append the terminal 0 themselves, so the returned list omits that 0.
|
| 81 |
+
|
| 82 |
+
Args:
|
| 83 |
+
num_inference_steps (`int`):
|
| 84 |
+
Number of Euler steps. Must be positive.
|
| 85 |
+
|
| 86 |
+
Returns:
|
| 87 |
+
`list[float]`: `num_inference_steps` sigmas starting at 1 and ending at `1 / steps`.
|
| 88 |
+
"""
|
| 89 |
+
if num_inference_steps <= 0:
|
| 90 |
+
raise ValueError(f"`num_inference_steps` must be positive, got {num_inference_steps}.")
|
| 91 |
+
flow_time = torch.linspace(0.0, 1.0, num_inference_steps + 1)
|
| 92 |
+
return (1.0 - flow_time)[:-1].tolist()
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
class LoopedDiTPipeline(DiffusionPipeline):
|
| 96 |
+
r"""
|
| 97 |
+
Text-to-image pipeline for Looped-DiT.
|
| 98 |
+
|
| 99 |
+
Looped-DiT denoises directly in RGB pixel space (no VAE). The transformer predicts the clean
|
| 100 |
+
image `x0`. This pipeline converts that prediction to a flow-matching velocity and integrates it
|
| 101 |
+
with a diffusers scheduler. The default scheduler is [`FlowMatchEulerDiscreteScheduler`] on the
|
| 102 |
+
same uniform grid as the paper (100 steps, shift 1). Any [`KarrasDiffusionSchedulers`] instance
|
| 103 |
+
can be assigned to `pipe.scheduler` without other code changes.
|
| 104 |
+
|
| 105 |
+
Classifier-free guidance uses the training null condition: an all-zero text mask, which the
|
| 106 |
+
denoiser replaces with its mask token. There is no separate negative-prompt encoder.
|
| 107 |
+
|
| 108 |
+
The pipeline inherits from [`DiffusionPipeline`]. Check the superclass documentation for the
|
| 109 |
+
generic methods (download, save, device placement, CPU offload).
|
| 110 |
+
|
| 111 |
+
Args:
|
| 112 |
+
transformer ([`ModelMixin`]):
|
| 113 |
+
Looped-DiT denoiser (`LoopedDiTTransformer2DModel`) that predicts `x0` in pixel space.
|
| 114 |
+
scheduler ([`FlowMatchEulerDiscreteScheduler`] or [`KarrasDiffusionSchedulers`]):
|
| 115 |
+
Scheduler used to step the flow. The paper setting is
|
| 116 |
+
`FlowMatchEulerDiscreteScheduler(num_train_timesteps=1000, shift=1.0)`.
|
| 117 |
+
tokenizer ([`~transformers.AutoTokenizer`], *optional*):
|
| 118 |
+
Tokenizer for the frozen text encoder. Loaded from `text_encoder_name` when missing.
|
| 119 |
+
text_encoder ([`~transformers.T5EncoderModel`], *optional*):
|
| 120 |
+
Frozen FLAN-T5 encoder. Loaded from `text_encoder_name` when missing.
|
| 121 |
+
`DiffusionPipeline.from_pretrained(..., torch_dtype=torch.bfloat16)` keeps this encoder in bf16 with the denoiser.
|
| 122 |
+
text_encoder_name (`str`, defaults to `"google/flan-t5-large"`):
|
| 123 |
+
Hub id or local path used when `tokenizer` / `text_encoder` are not passed.
|
| 124 |
+
prompt_length (`int`, *optional*):
|
| 125 |
+
Token length prompts are padded or truncated to. Defaults to `tokenizer.model_max_length`.
|
| 126 |
+
noise_scale (`float`, defaults to 2.0):
|
| 127 |
+
Standard deviation of the initial Gaussian, matching the training noise scale.
|
| 128 |
+
default_num_inference_steps (`int`, defaults to 100):
|
| 129 |
+
Step count used when `__call__` does not pass `num_inference_steps`.
|
| 130 |
+
"""
|
| 131 |
+
|
| 132 |
+
model_cpu_offload_seq = "text_encoder->transformer"
|
| 133 |
+
_optional_components = ["tokenizer", "text_encoder"]
|
| 134 |
+
_callback_tensor_inputs = ["latents", "prompt_embeds", "prompt_attention_mask"]
|
| 135 |
+
|
| 136 |
+
def __init__(
|
| 137 |
+
self,
|
| 138 |
+
transformer: ModelMixin,
|
| 139 |
+
scheduler: KarrasDiffusionSchedulers | SchedulerMixin,
|
| 140 |
+
tokenizer: Any | None = None,
|
| 141 |
+
text_encoder: T5EncoderModel | None = None,
|
| 142 |
+
text_encoder_name: str = "google/flan-t5-large",
|
| 143 |
+
prompt_length: int | None = None,
|
| 144 |
+
noise_scale: float = 2.0,
|
| 145 |
+
default_num_inference_steps: int = 100,
|
| 146 |
+
):
|
| 147 |
+
super().__init__()
|
| 148 |
+
if prompt_length is None and tokenizer is not None:
|
| 149 |
+
prompt_length = int(getattr(tokenizer, "model_max_length", 256))
|
| 150 |
+
if prompt_length is None:
|
| 151 |
+
prompt_length = 256
|
| 152 |
+
if scheduler is None:
|
| 153 |
+
scheduler = self._default_scheduler()
|
| 154 |
+
if noise_scale <= 0:
|
| 155 |
+
raise ValueError(f"`noise_scale` must be positive, got {noise_scale}.")
|
| 156 |
+
if prompt_length < 1:
|
| 157 |
+
raise ValueError(f"`prompt_length` must be positive, got {prompt_length}.")
|
| 158 |
+
if default_num_inference_steps < 1:
|
| 159 |
+
raise ValueError(f"`default_num_inference_steps` must be positive, got {default_num_inference_steps}.")
|
| 160 |
+
|
| 161 |
+
self.register_modules(
|
| 162 |
+
transformer=transformer,
|
| 163 |
+
scheduler=scheduler,
|
| 164 |
+
tokenizer=tokenizer,
|
| 165 |
+
text_encoder=text_encoder,
|
| 166 |
+
)
|
| 167 |
+
self.register_to_config(
|
| 168 |
+
text_encoder_name=text_encoder_name,
|
| 169 |
+
prompt_length=int(prompt_length),
|
| 170 |
+
noise_scale=float(noise_scale),
|
| 171 |
+
default_num_inference_steps=int(default_num_inference_steps),
|
| 172 |
+
)
|
| 173 |
+
|
| 174 |
+
@staticmethod
|
| 175 |
+
def _default_scheduler() -> FlowMatchEulerDiscreteScheduler:
|
| 176 |
+
r"""
|
| 177 |
+
Build the paper's Euler scheduler.
|
| 178 |
+
|
| 179 |
+
Returns:
|
| 180 |
+
[`FlowMatchEulerDiscreteScheduler`]: 1000 training timesteps, shift 1, deterministic.
|
| 181 |
+
"""
|
| 182 |
+
kwargs: dict[str, Any] = {"num_train_timesteps": 1000, "shift": 1.0}
|
| 183 |
+
if "stochastic_sampling" in inspect.signature(FlowMatchEulerDiscreteScheduler.__init__).parameters:
|
| 184 |
+
kwargs["stochastic_sampling"] = False
|
| 185 |
+
return FlowMatchEulerDiscreteScheduler(**kwargs)
|
| 186 |
+
|
| 187 |
+
def _encode_prompt(
|
| 188 |
+
self,
|
| 189 |
+
prompt: str | list[str] | None,
|
| 190 |
+
device: torch.device,
|
| 191 |
+
num_images_per_prompt: int,
|
| 192 |
+
prompt_embeds: torch.Tensor | None = None,
|
| 193 |
+
prompt_attention_mask: torch.Tensor | None = None,
|
| 194 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 195 |
+
r"""
|
| 196 |
+
Deprecated alias of [`~LoopedDiTPipeline.encode_prompt`].
|
| 197 |
+
|
| 198 |
+
Args:
|
| 199 |
+
prompt (`str` or `list[str]`, *optional*):
|
| 200 |
+
Prompt or prompts to tokenize and encode.
|
| 201 |
+
device (`torch.device`):
|
| 202 |
+
Device of the returned tensors.
|
| 203 |
+
num_images_per_prompt (`int`):
|
| 204 |
+
How many times to repeat each prompt embedding.
|
| 205 |
+
prompt_embeds (`torch.Tensor`, *optional*):
|
| 206 |
+
Already encoded prompts of shape `(batch, sequence, text_dim)`.
|
| 207 |
+
prompt_attention_mask (`torch.Tensor`, *optional*):
|
| 208 |
+
Mask of shape `(batch, sequence)` with 1 on real tokens. Required with `prompt_embeds`
|
| 209 |
+
only when padding should be replaced by the mask token; otherwise a mask of ones is used.
|
| 210 |
+
|
| 211 |
+
Returns:
|
| 212 |
+
`tuple[torch.Tensor, torch.Tensor]`: Prompt embeddings and the attention mask.
|
| 213 |
+
"""
|
| 214 |
+
deprecation_message = (
|
| 215 |
+
"`_encode_prompt()` is deprecated and will be removed in a future version. Use `encode_prompt()` instead."
|
| 216 |
+
)
|
| 217 |
+
deprecate("_encode_prompt()", "1.0.0", deprecation_message, standard_warn=False)
|
| 218 |
+
return self.encode_prompt(prompt, device, num_images_per_prompt, prompt_embeds, prompt_attention_mask)
|
| 219 |
+
|
| 220 |
+
def encode_prompt(
|
| 221 |
+
self,
|
| 222 |
+
prompt: str | list[str] | None,
|
| 223 |
+
device: torch.device,
|
| 224 |
+
num_images_per_prompt: int,
|
| 225 |
+
prompt_embeds: torch.Tensor | None = None,
|
| 226 |
+
prompt_attention_mask: torch.Tensor | None = None,
|
| 227 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 228 |
+
r"""
|
| 229 |
+
Encode prompts with the frozen FLAN-T5 encoder.
|
| 230 |
+
|
| 231 |
+
Prompts are padded or truncated to `config.prompt_length`. The unconditional branch of
|
| 232 |
+
classifier-free guidance is not encoded here: the denoiser builds it by zeroing this mask.
|
| 233 |
+
|
| 234 |
+
Args:
|
| 235 |
+
prompt (`str` or `list[str]`, *optional*):
|
| 236 |
+
Prompt or prompts to tokenize. Ignored when `prompt_embeds` is passed.
|
| 237 |
+
device (`torch.device`):
|
| 238 |
+
Device of the returned tensors.
|
| 239 |
+
num_images_per_prompt (`int`):
|
| 240 |
+
Number of times to repeat each encoded prompt along the batch dimension.
|
| 241 |
+
prompt_embeds (`torch.Tensor`, *optional*):
|
| 242 |
+
Precomputed embeddings of shape `(batch, sequence, text_dim)`. When set, `prompt` is ignored.
|
| 243 |
+
prompt_attention_mask (`torch.Tensor`, *optional*):
|
| 244 |
+
Mask of shape `(batch, sequence)`, 1 for tokens that should condition the model. When
|
| 245 |
+
`prompt_embeds` is set and this is omitted, every position is treated as a real token.
|
| 246 |
+
|
| 247 |
+
Returns:
|
| 248 |
+
`tuple[torch.Tensor, torch.Tensor]`:
|
| 249 |
+
Embeddings `(batch * num_images_per_prompt, sequence, text_dim)` and a mask of the same batch.
|
| 250 |
+
"""
|
| 251 |
+
if num_images_per_prompt < 1:
|
| 252 |
+
raise ValueError(f"`num_images_per_prompt` must be >= 1, got {num_images_per_prompt}.")
|
| 253 |
+
|
| 254 |
+
if prompt_embeds is None:
|
| 255 |
+
if isinstance(prompt, str):
|
| 256 |
+
prompt = [prompt]
|
| 257 |
+
if self.tokenizer is None:
|
| 258 |
+
self.tokenizer = AutoTokenizer.from_pretrained(
|
| 259 |
+
self.config.text_encoder_name, model_max_length=int(self.config.prompt_length)
|
| 260 |
+
)
|
| 261 |
+
if self.text_encoder is None:
|
| 262 |
+
self.text_encoder = T5EncoderModel.from_pretrained(self.config.text_encoder_name)
|
| 263 |
+
self.text_encoder.requires_grad_(False)
|
| 264 |
+
self.text_encoder.eval()
|
| 265 |
+
encoder_device = next(self.text_encoder.parameters()).device
|
| 266 |
+
if encoder_device != device:
|
| 267 |
+
self.text_encoder.to(device)
|
| 268 |
+
tokens = self.tokenizer(
|
| 269 |
+
prompt,
|
| 270 |
+
max_length=int(self.config.prompt_length),
|
| 271 |
+
padding="max_length",
|
| 272 |
+
truncation=True,
|
| 273 |
+
return_tensors="pt",
|
| 274 |
+
)
|
| 275 |
+
input_ids = tokens.input_ids.to(device)
|
| 276 |
+
prompt_attention_mask = tokens.attention_mask.to(device)
|
| 277 |
+
prompt_embeds = self.text_encoder(input_ids=input_ids, attention_mask=prompt_attention_mask).last_hidden_state
|
| 278 |
+
else:
|
| 279 |
+
prompt_embeds = prompt_embeds.to(device)
|
| 280 |
+
if prompt_attention_mask is None:
|
| 281 |
+
prompt_attention_mask = torch.ones(
|
| 282 |
+
prompt_embeds.shape[:2], device=device, dtype=torch.long
|
| 283 |
+
)
|
| 284 |
+
else:
|
| 285 |
+
prompt_attention_mask = prompt_attention_mask.to(device)
|
| 286 |
+
if prompt_embeds.shape[0] != prompt_attention_mask.shape[0]:
|
| 287 |
+
raise ValueError(
|
| 288 |
+
"`prompt_embeds` and `prompt_attention_mask` must have the same batch size, got "
|
| 289 |
+
f"{prompt_embeds.shape[0]} and {prompt_attention_mask.shape[0]}."
|
| 290 |
+
)
|
| 291 |
+
|
| 292 |
+
if num_images_per_prompt != 1:
|
| 293 |
+
prompt_embeds = prompt_embeds.repeat_interleave(num_images_per_prompt, dim=0)
|
| 294 |
+
prompt_attention_mask = prompt_attention_mask.repeat_interleave(num_images_per_prompt, dim=0)
|
| 295 |
+
return prompt_embeds, prompt_attention_mask
|
| 296 |
+
|
| 297 |
+
def prepare_extra_step_kwargs(
|
| 298 |
+
self, generator: torch.Generator | list[torch.Generator] | None, eta: float
|
| 299 |
+
) -> dict[str, Any]:
|
| 300 |
+
r"""
|
| 301 |
+
Extra arguments forwarded to `scheduler.step`, depending on what that method accepts.
|
| 302 |
+
|
| 303 |
+
Args:
|
| 304 |
+
generator (`torch.Generator` or `list[torch.Generator]`, *optional*):
|
| 305 |
+
Generator passed through when the scheduler step samples noise.
|
| 306 |
+
eta (`float`):
|
| 307 |
+
DDIM eta in `[0, 1]`. Ignored by schedulers whose `step` has no `eta` argument.
|
| 308 |
+
|
| 309 |
+
Returns:
|
| 310 |
+
`dict`: Keyword arguments for `scheduler.step`.
|
| 311 |
+
"""
|
| 312 |
+
extra_step_kwargs: dict[str, Any] = {}
|
| 313 |
+
step_params = set(inspect.signature(self.scheduler.step).parameters.keys())
|
| 314 |
+
if "eta" in step_params:
|
| 315 |
+
extra_step_kwargs["eta"] = eta
|
| 316 |
+
if "generator" in step_params:
|
| 317 |
+
extra_step_kwargs["generator"] = generator
|
| 318 |
+
return extra_step_kwargs
|
| 319 |
+
|
| 320 |
+
def check_inputs(
|
| 321 |
+
self,
|
| 322 |
+
prompt: str | list[str] | None,
|
| 323 |
+
height: int,
|
| 324 |
+
width: int,
|
| 325 |
+
callback_steps: int | None,
|
| 326 |
+
prompt_embeds: torch.Tensor | None = None,
|
| 327 |
+
prompt_attention_mask: torch.Tensor | None = None,
|
| 328 |
+
callback_on_step_end_tensor_inputs: list[str] | None = None,
|
| 329 |
+
num_inference_steps: int = 100,
|
| 330 |
+
guidance_scale: float = 6.0,
|
| 331 |
+
num_loops: int | None = None,
|
| 332 |
+
output_type: str = "pil",
|
| 333 |
+
) -> None:
|
| 334 |
+
r"""
|
| 335 |
+
Validate generation arguments and raise `ValueError` or `TypeError` on misuse.
|
| 336 |
+
|
| 337 |
+
Args:
|
| 338 |
+
prompt (`str` or `list[str]`, *optional*):
|
| 339 |
+
Prompt text. Mutually exclusive with `prompt_embeds`.
|
| 340 |
+
height (`int`):
|
| 341 |
+
Output height in pixels. Must equal the transformer's trained `image_size`.
|
| 342 |
+
width (`int`):
|
| 343 |
+
Output width in pixels. Must equal the transformer's trained `image_size`.
|
| 344 |
+
callback_steps (`int`, *optional*):
|
| 345 |
+
Deprecated callback period. When set, it must be a positive integer.
|
| 346 |
+
prompt_embeds (`torch.Tensor`, *optional*):
|
| 347 |
+
Precomputed text embeddings. Required when `prompt` is omitted.
|
| 348 |
+
prompt_attention_mask (`torch.Tensor`, *optional*):
|
| 349 |
+
Mask paired with `prompt_embeds`.
|
| 350 |
+
callback_on_step_end_tensor_inputs (`list[str]`, *optional*):
|
| 351 |
+
Tensor names the step callback may read. Each name must be listed on
|
| 352 |
+
`_callback_tensor_inputs`.
|
| 353 |
+
num_inference_steps (`int`):
|
| 354 |
+
Denoising steps. Must be positive.
|
| 355 |
+
guidance_scale (`float`):
|
| 356 |
+
Classifier-free guidance scale. Must be finite. `1` disables guidance.
|
| 357 |
+
num_loops (`int`, *optional*):
|
| 358 |
+
Loop depth. `None` uses the depth stored on the transformer. Otherwise `>= 1`, and
|
| 359 |
+
untied models cannot exceed the trained depth.
|
| 360 |
+
output_type (`str`):
|
| 361 |
+
One of `"pil"`, `"np"`, `"pt"`, or `"latent"`.
|
| 362 |
+
"""
|
| 363 |
+
image_size = int(self.transformer.config.image_size)
|
| 364 |
+
patch_size = int(self.transformer.config.patch_size)
|
| 365 |
+
if height != image_size or width != image_size:
|
| 366 |
+
raise ValueError(
|
| 367 |
+
f"Looped-DiT uses a fixed positional grid of {image_size}x{image_size} "
|
| 368 |
+
f"(patch size {patch_size}). Got height={height}, width={width}."
|
| 369 |
+
)
|
| 370 |
+
if height % patch_size != 0 or width % patch_size != 0:
|
| 371 |
+
raise ValueError(f"height and width must be divisible by patch_size={patch_size}, got {(height, width)}.")
|
| 372 |
+
|
| 373 |
+
if callback_steps is not None and (not isinstance(callback_steps, int) or callback_steps <= 0):
|
| 374 |
+
raise ValueError(
|
| 375 |
+
f"`callback_steps` has to be a positive integer but is {callback_steps} of type {type(callback_steps)}."
|
| 376 |
+
)
|
| 377 |
+
if callback_on_step_end_tensor_inputs is not None and not all(
|
| 378 |
+
key in self._callback_tensor_inputs for key in callback_on_step_end_tensor_inputs
|
| 379 |
+
):
|
| 380 |
+
unexpected = [key for key in callback_on_step_end_tensor_inputs if key not in self._callback_tensor_inputs]
|
| 381 |
+
raise ValueError(
|
| 382 |
+
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {unexpected}."
|
| 383 |
+
)
|
| 384 |
+
|
| 385 |
+
if prompt is not None and prompt_embeds is not None:
|
| 386 |
+
raise ValueError("Cannot forward both `prompt` and `prompt_embeds`. Pass only one of them.")
|
| 387 |
+
if prompt is None and prompt_embeds is None:
|
| 388 |
+
raise ValueError("Provide either `prompt` or `prompt_embeds`.")
|
| 389 |
+
if prompt is not None and not isinstance(prompt, str) and not (
|
| 390 |
+
isinstance(prompt, list) and all(isinstance(item, str) for item in prompt)
|
| 391 |
+
):
|
| 392 |
+
raise TypeError(f"`prompt` has to be a string or a list of strings, got {type(prompt)}.")
|
| 393 |
+
if prompt_embeds is not None and prompt_embeds.ndim != 3:
|
| 394 |
+
raise ValueError(f"`prompt_embeds` must have shape (batch, sequence, dim), got {tuple(prompt_embeds.shape)}.")
|
| 395 |
+
if prompt_attention_mask is not None and prompt_embeds is None:
|
| 396 |
+
raise ValueError("`prompt_attention_mask` was passed without `prompt_embeds`.")
|
| 397 |
+
|
| 398 |
+
if num_inference_steps <= 0:
|
| 399 |
+
raise ValueError(f"`num_inference_steps` must be positive, got {num_inference_steps}.")
|
| 400 |
+
if not torch.isfinite(torch.tensor(guidance_scale)):
|
| 401 |
+
raise ValueError(f"`guidance_scale` must be finite, got {guidance_scale}.")
|
| 402 |
+
if num_loops is not None:
|
| 403 |
+
if int(num_loops) < 1:
|
| 404 |
+
raise ValueError(f"`num_loops` must be >= 1, got {num_loops}.")
|
| 405 |
+
trained = int(self.transformer.config.num_loops)
|
| 406 |
+
if not bool(self.transformer.config.share_loop_weights) and int(num_loops) > trained:
|
| 407 |
+
raise ValueError(
|
| 408 |
+
f"This checkpoint does not share loop weights, so `num_loops` cannot exceed the trained "
|
| 409 |
+
f"depth {trained}. Got {num_loops}."
|
| 410 |
+
)
|
| 411 |
+
if output_type not in {"pil", "np", "pt", "latent"}:
|
| 412 |
+
raise ValueError(f"Unsupported `output_type` {output_type!r}. Choose from 'pil', 'np', 'pt', 'latent'.")
|
| 413 |
+
|
| 414 |
+
def prepare_latents(
|
| 415 |
+
self,
|
| 416 |
+
batch_size: int,
|
| 417 |
+
num_channels: int,
|
| 418 |
+
height: int,
|
| 419 |
+
width: int,
|
| 420 |
+
dtype: torch.dtype,
|
| 421 |
+
device: torch.device,
|
| 422 |
+
generator: torch.Generator | list[torch.Generator] | None,
|
| 423 |
+
latents: torch.Tensor | None = None,
|
| 424 |
+
) -> torch.Tensor:
|
| 425 |
+
r"""
|
| 426 |
+
Sample the initial pixel-space noise, or validate a tensor the caller already sampled.
|
| 427 |
+
|
| 428 |
+
Looped-DiT has no VAE, so "latents" here are RGB images. Fresh noise is scaled by
|
| 429 |
+
`config.noise_scale` (2.0 in the paper). A provided `latents` tensor is not rescaled.
|
| 430 |
+
|
| 431 |
+
Args:
|
| 432 |
+
batch_size (`int`):
|
| 433 |
+
Number of images, including `num_images_per_prompt`.
|
| 434 |
+
num_channels (`int`):
|
| 435 |
+
Channel count. 3 for RGB.
|
| 436 |
+
height (`int`):
|
| 437 |
+
Image height in pixels.
|
| 438 |
+
width (`int`):
|
| 439 |
+
Image width in pixels.
|
| 440 |
+
dtype (`torch.dtype`):
|
| 441 |
+
Dtype of freshly sampled noise. The integration itself is accumulated in float32.
|
| 442 |
+
device (`torch.device`):
|
| 443 |
+
Device of the returned tensor.
|
| 444 |
+
generator (`torch.Generator` or `list[torch.Generator]`, *optional*):
|
| 445 |
+
Per-call RNG. A list must have length `batch_size`.
|
| 446 |
+
latents (`torch.Tensor`, *optional*):
|
| 447 |
+
Starting noise of shape `(batch_size, num_channels, height, width)`.
|
| 448 |
+
|
| 449 |
+
Returns:
|
| 450 |
+
`torch.Tensor`: Starting noise of shape `(batch_size, num_channels, height, width)`.
|
| 451 |
+
"""
|
| 452 |
+
shape = (batch_size, num_channels, height, width)
|
| 453 |
+
if isinstance(generator, list) and len(generator) != batch_size:
|
| 454 |
+
raise ValueError(
|
| 455 |
+
f"You passed a list of {len(generator)} generators for a batch of {batch_size}. "
|
| 456 |
+
"The two lengths must match."
|
| 457 |
+
)
|
| 458 |
+
if latents is None:
|
| 459 |
+
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
| 460 |
+
latents = latents * float(self.config.noise_scale)
|
| 461 |
+
else:
|
| 462 |
+
latents = latents.to(device=device)
|
| 463 |
+
if tuple(latents.shape) != shape:
|
| 464 |
+
raise ValueError(f"`latents` shape {tuple(latents.shape)} does not match the expected {shape}.")
|
| 465 |
+
return latents
|
| 466 |
+
|
| 467 |
+
@property
|
| 468 |
+
def guidance_scale(self) -> float:
|
| 469 |
+
r"""
|
| 470 |
+
Classifier-free guidance scale of the call that is currently running.
|
| 471 |
+
|
| 472 |
+
Returns:
|
| 473 |
+
`float`: The scale set by the active `__call__`.
|
| 474 |
+
"""
|
| 475 |
+
return self._guidance_scale
|
| 476 |
+
|
| 477 |
+
@property
|
| 478 |
+
def do_classifier_free_guidance(self) -> bool:
|
| 479 |
+
r"""
|
| 480 |
+
Whether the active call runs a conditional and an unconditional forward.
|
| 481 |
+
|
| 482 |
+
Returns:
|
| 483 |
+
`bool`: True when `guidance_scale != 1`.
|
| 484 |
+
"""
|
| 485 |
+
return self._guidance_scale != 1.0
|
| 486 |
+
|
| 487 |
+
@property
|
| 488 |
+
def num_timesteps(self) -> int:
|
| 489 |
+
r"""
|
| 490 |
+
Number of scheduler timesteps in the active call.
|
| 491 |
+
|
| 492 |
+
Returns:
|
| 493 |
+
`int`: Length of the timestep schedule.
|
| 494 |
+
"""
|
| 495 |
+
return self._num_timesteps
|
| 496 |
+
|
| 497 |
+
@property
|
| 498 |
+
def interrupt(self) -> bool:
|
| 499 |
+
r"""
|
| 500 |
+
Whether the active denoising loop should skip remaining steps.
|
| 501 |
+
|
| 502 |
+
Returns:
|
| 503 |
+
`bool`: True after the caller sets `pipeline._interrupt = True`.
|
| 504 |
+
"""
|
| 505 |
+
return self._interrupt
|
| 506 |
+
|
| 507 |
+
def _images_from_latents(self, latents: torch.Tensor, output_type: str) -> torch.Tensor | list[Image.Image] | Any:
|
| 508 |
+
r"""
|
| 509 |
+
Convert pixel-space samples in `[-1, 1]` to the requested output type.
|
| 510 |
+
|
| 511 |
+
Quantization matches the original sampler: `uint8(clamp(x, -1, 1) * 127.5 + 128)`.
|
| 512 |
+
|
| 513 |
+
Args:
|
| 514 |
+
latents (`torch.Tensor`):
|
| 515 |
+
Samples of shape `(batch, channels, height, width)` in model range `[-1, 1]`.
|
| 516 |
+
output_type (`str`):
|
| 517 |
+
`"latent"` returns `latents` unchanged. `"pt"` is float RGB in `[0, 1]`. `"np"` is
|
| 518 |
+
`uint8` HWC arrays. `"pil"` is a list of `PIL.Image.Image`.
|
| 519 |
+
|
| 520 |
+
Returns:
|
| 521 |
+
Images in the requested type.
|
| 522 |
+
"""
|
| 523 |
+
if output_type == "latent":
|
| 524 |
+
return latents
|
| 525 |
+
images = (latents.float().clamp(-1, 1) * 127.5 + 128.0).clamp(0, 255).to(torch.uint8)
|
| 526 |
+
if output_type == "pt":
|
| 527 |
+
return images.float() / 255.0
|
| 528 |
+
arrays = images.permute(0, 2, 3, 1).cpu().numpy()
|
| 529 |
+
if output_type == "np":
|
| 530 |
+
return arrays
|
| 531 |
+
return [Image.fromarray(image) for image in arrays]
|
| 532 |
+
|
| 533 |
+
@torch.no_grad()
|
| 534 |
+
@replace_example_docstring(EXAMPLE_DOC_STRING)
|
| 535 |
+
def __call__(
|
| 536 |
+
self,
|
| 537 |
+
prompt: str | list[str] | None = None,
|
| 538 |
+
height: int | None = None,
|
| 539 |
+
width: int | None = None,
|
| 540 |
+
num_inference_steps: int | None = None,
|
| 541 |
+
timesteps: list[int] | None = None,
|
| 542 |
+
sigmas: list[float] | None = None,
|
| 543 |
+
guidance_scale: float = 6.0,
|
| 544 |
+
num_images_per_prompt: int = 1,
|
| 545 |
+
num_loops: int | None = None,
|
| 546 |
+
eta: float = 0.0,
|
| 547 |
+
generator: torch.Generator | list[torch.Generator] | None = None,
|
| 548 |
+
latents: torch.Tensor | None = None,
|
| 549 |
+
prompt_embeds: torch.Tensor | None = None,
|
| 550 |
+
prompt_attention_mask: torch.Tensor | None = None,
|
| 551 |
+
output_type: str = "pil",
|
| 552 |
+
return_dict: bool = True,
|
| 553 |
+
callback_on_step_end: Callable[[int, int, dict], dict] | PipelineCallback | MultiPipelineCallbacks | None = None,
|
| 554 |
+
callback_on_step_end_tensor_inputs: list[str] = ["latents"],
|
| 555 |
+
**kwargs,
|
| 556 |
+
) -> ImagePipelineOutput | tuple:
|
| 557 |
+
r"""
|
| 558 |
+
Generate images from text prompts.
|
| 559 |
+
|
| 560 |
+
Args:
|
| 561 |
+
prompt (`str` or `list[str]`, *optional*):
|
| 562 |
+
Prompt or prompts to guide image generation. Required unless `prompt_embeds` is passed.
|
| 563 |
+
height (`int`, *optional*):
|
| 564 |
+
Image height in pixels. Defaults to the transformer's trained resolution (512).
|
| 565 |
+
Other resolutions are rejected: the positional embedding is a fixed grid.
|
| 566 |
+
width (`int`, *optional*):
|
| 567 |
+
Image width in pixels. Defaults to the trained resolution and must match `height`.
|
| 568 |
+
num_inference_steps (`int`, *optional*):
|
| 569 |
+
Denoising steps. Defaults to `config.default_num_inference_steps` (100).
|
| 570 |
+
timesteps (`list[int]`, *optional*):
|
| 571 |
+
Custom scheduler timesteps, descending. Mutually exclusive with `sigmas`. Ignored by the
|
| 572 |
+
paper Euler grid, which is selected only when both `timesteps` and `sigmas` are omitted
|
| 573 |
+
and the scheduler is [`FlowMatchEulerDiscreteScheduler`].
|
| 574 |
+
sigmas (`list[float]`, *optional*):
|
| 575 |
+
Custom sigmas passed to `scheduler.set_timesteps`. Mutually exclusive with `timesteps`.
|
| 576 |
+
guidance_scale (`float`, defaults to 6.0):
|
| 577 |
+
Classifier-free guidance scale from the paper. Guidance is on when this is not `1`.
|
| 578 |
+
The unconditional branch is an empty text mask, not a negative prompt.
|
| 579 |
+
num_images_per_prompt (`int`, defaults to 1):
|
| 580 |
+
How many images to sample for each prompt.
|
| 581 |
+
num_loops (`int`, *optional*):
|
| 582 |
+
How many times to run the shared middle blocks. `None` uses the trained depth. Other
|
| 583 |
+
depths work without retraining when loop weights are shared.
|
| 584 |
+
eta (`float`, defaults to 0.0):
|
| 585 |
+
DDIM eta. Ignored by the flow-match Euler scheduler.
|
| 586 |
+
generator (`torch.Generator` or `list[torch.Generator]`, *optional*):
|
| 587 |
+
RNG for the initial noise. `None` uses PyTorch's global generator, which is what
|
| 588 |
+
`torch.manual_seed` seeds.
|
| 589 |
+
latents (`torch.Tensor`, *optional*):
|
| 590 |
+
Initial noise `(batch, 3, height, width)`. Not multiplied by `noise_scale`.
|
| 591 |
+
prompt_embeds (`torch.Tensor`, *optional*):
|
| 592 |
+
Precomputed FLAN-T5 states `(batch, sequence, text_dim)` in place of `prompt`.
|
| 593 |
+
prompt_attention_mask (`torch.Tensor`, *optional*):
|
| 594 |
+
Mask `(batch, sequence)` paired with `prompt_embeds`. 1 marks real tokens.
|
| 595 |
+
output_type (`str`, defaults to `"pil"`):
|
| 596 |
+
`"pil"`, `"np"`, `"pt"` (float RGB in `[0, 1]`), or `"latent"` (pixels in model range).
|
| 597 |
+
return_dict (`bool`, defaults to `True`):
|
| 598 |
+
Return [`ImagePipelineOutput`] when `True`, otherwise a one-tuple of images.
|
| 599 |
+
callback_on_step_end (`Callable` or `PipelineCallback`, *optional*):
|
| 600 |
+
Called as `callback_on_step_end(pipeline, step, timestep, callback_kwargs)` after each
|
| 601 |
+
scheduler step. Return a dict to replace tensors listed in
|
| 602 |
+
`callback_on_step_end_tensor_inputs`.
|
| 603 |
+
callback_on_step_end_tensor_inputs (`list[str]`, defaults to `["latents"]`):
|
| 604 |
+
Tensor names passed to the step callback. Must be a subset of `_callback_tensor_inputs`.
|
| 605 |
+
|
| 606 |
+
Examples:
|
| 607 |
+
|
| 608 |
+
Returns:
|
| 609 |
+
[`ImagePipelineOutput`] or `tuple`:
|
| 610 |
+
When `return_dict` is `True`, [`ImagePipelineOutput`] with the images. Otherwise a tuple
|
| 611 |
+
whose first element is the images.
|
| 612 |
+
"""
|
| 613 |
+
callback = kwargs.pop("callback", None)
|
| 614 |
+
callback_steps = kwargs.pop("callback_steps", None)
|
| 615 |
+
if kwargs:
|
| 616 |
+
raise TypeError(f"Unexpected arguments: {sorted(kwargs)}.")
|
| 617 |
+
|
| 618 |
+
if callback is not None:
|
| 619 |
+
deprecate(
|
| 620 |
+
"callback",
|
| 621 |
+
"1.0.0",
|
| 622 |
+
"Passing `callback` as an input argument to `__call__` is deprecated, consider using `callback_on_step_end`",
|
| 623 |
+
)
|
| 624 |
+
if callback_steps is not None:
|
| 625 |
+
deprecate(
|
| 626 |
+
"callback_steps",
|
| 627 |
+
"1.0.0",
|
| 628 |
+
"Passing `callback_steps` as an input argument to `__call__` is deprecated, consider using `callback_on_step_end`",
|
| 629 |
+
)
|
| 630 |
+
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
|
| 631 |
+
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
|
| 632 |
+
|
| 633 |
+
image_size = int(self.transformer.config.image_size)
|
| 634 |
+
height = image_size if height is None else int(height)
|
| 635 |
+
width = image_size if width is None else int(width)
|
| 636 |
+
if num_inference_steps is None:
|
| 637 |
+
num_inference_steps = int(self.config.default_num_inference_steps)
|
| 638 |
+
|
| 639 |
+
# 1. Check inputs.
|
| 640 |
+
self.check_inputs(
|
| 641 |
+
prompt,
|
| 642 |
+
height,
|
| 643 |
+
width,
|
| 644 |
+
callback_steps,
|
| 645 |
+
prompt_embeds,
|
| 646 |
+
prompt_attention_mask,
|
| 647 |
+
callback_on_step_end_tensor_inputs,
|
| 648 |
+
num_inference_steps,
|
| 649 |
+
guidance_scale,
|
| 650 |
+
num_loops,
|
| 651 |
+
output_type,
|
| 652 |
+
)
|
| 653 |
+
|
| 654 |
+
self._guidance_scale = float(guidance_scale)
|
| 655 |
+
self._interrupt = False
|
| 656 |
+
|
| 657 |
+
# 2. Define call parameters.
|
| 658 |
+
if prompt is not None and isinstance(prompt, str):
|
| 659 |
+
batch_size = 1
|
| 660 |
+
elif prompt is not None and isinstance(prompt, list):
|
| 661 |
+
batch_size = len(prompt)
|
| 662 |
+
else:
|
| 663 |
+
batch_size = prompt_embeds.shape[0]
|
| 664 |
+
device = self._execution_device
|
| 665 |
+
|
| 666 |
+
# 3. Encode input prompt.
|
| 667 |
+
prompt_embeds, prompt_attention_mask = self.encode_prompt(
|
| 668 |
+
prompt,
|
| 669 |
+
device,
|
| 670 |
+
num_images_per_prompt,
|
| 671 |
+
prompt_embeds=prompt_embeds,
|
| 672 |
+
prompt_attention_mask=prompt_attention_mask,
|
| 673 |
+
)
|
| 674 |
+
prompt_embeds = prompt_embeds.to(device=device, dtype=self.transformer.dtype)
|
| 675 |
+
prompt_attention_mask = prompt_attention_mask.to(device=device)
|
| 676 |
+
|
| 677 |
+
# 4. Prepare timesteps.
|
| 678 |
+
# The paper's Euler grid is the training linspace. Other schedulers keep their own spacing.
|
| 679 |
+
# A caller-supplied `timesteps` or `sigmas` always wins.
|
| 680 |
+
if timesteps is None and sigmas is None and isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler):
|
| 681 |
+
sigmas = paper_euler_sigmas(num_inference_steps)
|
| 682 |
+
timesteps, num_inference_steps = retrieve_timesteps(self.scheduler, num_inference_steps, device, timesteps, sigmas)
|
| 683 |
+
if getattr(self.scheduler.config, "stochastic_sampling", False):
|
| 684 |
+
raise ValueError(
|
| 685 |
+
"Looped-DiT's training sampler is deterministic. Set `stochastic_sampling=False` on "
|
| 686 |
+
"FlowMatchEulerDiscreteScheduler, or assign a different scheduler."
|
| 687 |
+
)
|
| 688 |
+
|
| 689 |
+
# 5. Prepare latent variables (pixel-space noise; there is no VAE).
|
| 690 |
+
latents = self.prepare_latents(
|
| 691 |
+
batch_size * num_images_per_prompt,
|
| 692 |
+
int(self.transformer.config.in_channels),
|
| 693 |
+
height,
|
| 694 |
+
width,
|
| 695 |
+
self.transformer.dtype,
|
| 696 |
+
device,
|
| 697 |
+
generator,
|
| 698 |
+
latents,
|
| 699 |
+
)
|
| 700 |
+
|
| 701 |
+
# 6. Prepare extra step kwargs.
|
| 702 |
+
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
|
| 703 |
+
num_train_timesteps = int(self.scheduler.config.num_train_timesteps)
|
| 704 |
+
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
| 705 |
+
self._num_timesteps = len(timesteps)
|
| 706 |
+
|
| 707 |
+
# 7. Denoising loop.
|
| 708 |
+
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
| 709 |
+
for i, t in enumerate(timesteps):
|
| 710 |
+
if self.interrupt:
|
| 711 |
+
continue
|
| 712 |
+
|
| 713 |
+
model_latents = latents
|
| 714 |
+
if hasattr(self.scheduler, "scale_model_input"):
|
| 715 |
+
model_latents = self.scheduler.scale_model_input(model_latents, t)
|
| 716 |
+
text = prompt_embeds
|
| 717 |
+
mask = prompt_attention_mask
|
| 718 |
+
if self.do_classifier_free_guidance:
|
| 719 |
+
model_latents = torch.cat([model_latents, model_latents], dim=0)
|
| 720 |
+
text = torch.cat([text, text], dim=0)
|
| 721 |
+
mask = torch.cat([mask, torch.zeros_like(mask)], dim=0)
|
| 722 |
+
|
| 723 |
+
# fp32 latents with bf16 weights match the old sampler, which autocasts the forward.
|
| 724 |
+
amp_dtype = self.transformer.dtype
|
| 725 |
+
use_amp = model_latents.is_cuda and amp_dtype in (torch.float16, torch.bfloat16)
|
| 726 |
+
if not use_amp:
|
| 727 |
+
model_latents = model_latents.to(dtype=amp_dtype)
|
| 728 |
+
with torch.autocast("cuda", dtype=amp_dtype, enabled=use_amp):
|
| 729 |
+
x0 = self.transformer(model_latents, text, mask, num_loops=num_loops)
|
| 730 |
+
x0 = x0.float()
|
| 731 |
+
if self.do_classifier_free_guidance:
|
| 732 |
+
x0_cond, x0_uncond = x0.chunk(2)
|
| 733 |
+
x0 = x0_uncond + self.guidance_scale * (x0_cond - x0_uncond)
|
| 734 |
+
|
| 735 |
+
# sigma = 1 - t_flow. Passing -velocity makes `x + (sigma_next - sigma) * model_output`
|
| 736 |
+
# equal the training update `x + velocity * (t_next - t)`.
|
| 737 |
+
flow_time = (1.0 - t.to(device=latents.device, dtype=torch.float32) / num_train_timesteps)
|
| 738 |
+
velocity = (x0 - latents.float()) / (1.0 - flow_time).clamp_min(VELOCITY_DENOM_MIN)
|
| 739 |
+
latents = self.scheduler.step(-velocity, t, latents, **extra_step_kwargs, return_dict=False)[0]
|
| 740 |
+
|
| 741 |
+
if callback_on_step_end is not None:
|
| 742 |
+
callback_kwargs = {key: locals()[key] for key in callback_on_step_end_tensor_inputs}
|
| 743 |
+
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
|
| 744 |
+
latents = callback_outputs.pop("latents", latents)
|
| 745 |
+
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
|
| 746 |
+
prompt_attention_mask = callback_outputs.pop("prompt_attention_mask", prompt_attention_mask)
|
| 747 |
+
|
| 748 |
+
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
| 749 |
+
progress_bar.update()
|
| 750 |
+
if callback is not None and i % callback_steps == 0:
|
| 751 |
+
callback(i, t, latents)
|
| 752 |
+
|
| 753 |
+
if XLA_AVAILABLE:
|
| 754 |
+
xm.mark_step()
|
| 755 |
+
|
| 756 |
+
images = self._images_from_latents(latents, output_type)
|
| 757 |
+
|
| 758 |
+
# Offload all models.
|
| 759 |
+
self.maybe_free_model_hooks()
|
| 760 |
+
|
| 761 |
+
if not return_dict:
|
| 762 |
+
return (images,)
|
| 763 |
+
return ImagePipelineOutput(images=images)
|
Looped-DiT-B-16/scheduler/scheduler_config.json
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"_class_name": "FlowMatchEulerDiscreteScheduler",
|
| 3 |
+
"_diffusers_version": "0.39.0",
|
| 4 |
+
"base_image_seq_len": 256,
|
| 5 |
+
"base_shift": 0.5,
|
| 6 |
+
"invert_sigmas": false,
|
| 7 |
+
"max_image_seq_len": 4096,
|
| 8 |
+
"max_shift": 1.15,
|
| 9 |
+
"num_train_timesteps": 1000,
|
| 10 |
+
"shift": 1.0,
|
| 11 |
+
"shift_terminal": null,
|
| 12 |
+
"stochastic_sampling": false,
|
| 13 |
+
"time_shift_type": "exponential",
|
| 14 |
+
"use_beta_sigmas": false,
|
| 15 |
+
"use_dynamic_shifting": false,
|
| 16 |
+
"use_exponential_sigmas": false,
|
| 17 |
+
"use_karras_sigmas": false
|
| 18 |
+
}
|
Looped-DiT-B-16/text_encoder/README.md
ADDED
|
@@ -0,0 +1,276 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
language:
|
| 3 |
+
- en
|
| 4 |
+
- fr
|
| 5 |
+
- ro
|
| 6 |
+
- de
|
| 7 |
+
- multilingual
|
| 8 |
+
|
| 9 |
+
widget:
|
| 10 |
+
- text: "Translate to German: My name is Arthur"
|
| 11 |
+
example_title: "Translation"
|
| 12 |
+
- text: "Please answer to the following question. Who is going to be the next Ballon d'or?"
|
| 13 |
+
example_title: "Question Answering"
|
| 14 |
+
- text: "Q: Can Geoffrey Hinton have a conversation with George Washington? Give the rationale before answering."
|
| 15 |
+
example_title: "Logical reasoning"
|
| 16 |
+
- text: "Please answer the following question. What is the boiling point of Nitrogen?"
|
| 17 |
+
example_title: "Scientific knowledge"
|
| 18 |
+
- text: "Answer the following yes/no question. Can you write a whole Haiku in a single tweet?"
|
| 19 |
+
example_title: "Yes/no question"
|
| 20 |
+
- text: "Answer the following yes/no question by reasoning step-by-step. Can you write a whole Haiku in a single tweet?"
|
| 21 |
+
example_title: "Reasoning task"
|
| 22 |
+
- text: "Q: ( False or not False or False ) is? A: Let's think step by step"
|
| 23 |
+
example_title: "Boolean Expressions"
|
| 24 |
+
- text: "The square root of x is the cube root of y. What is y to the power of 2, if x = 4?"
|
| 25 |
+
example_title: "Math reasoning"
|
| 26 |
+
- text: "Premise: At my age you will probably have learnt one lesson. Hypothesis: It's not certain how many lessons you'll learn by your thirties. Does the premise entail the hypothesis?"
|
| 27 |
+
example_title: "Premise and hypothesis"
|
| 28 |
+
|
| 29 |
+
tags:
|
| 30 |
+
- text2text-generation
|
| 31 |
+
|
| 32 |
+
datasets:
|
| 33 |
+
- svakulenk0/qrecc
|
| 34 |
+
- taskmaster2
|
| 35 |
+
- djaym7/wiki_dialog
|
| 36 |
+
- deepmind/code_contests
|
| 37 |
+
- lambada
|
| 38 |
+
- gsm8k
|
| 39 |
+
- aqua_rat
|
| 40 |
+
- esnli
|
| 41 |
+
- quasc
|
| 42 |
+
- qed
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
license: apache-2.0
|
| 46 |
+
---
|
| 47 |
+
|
| 48 |
+
# Model Card for FLAN-T5 large
|
| 49 |
+
|
| 50 |
+
<img src="https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/transformers/model_doc/flan2_architecture.jpg"
|
| 51 |
+
alt="drawing" width="600"/>
|
| 52 |
+
|
| 53 |
+
# Table of Contents
|
| 54 |
+
|
| 55 |
+
0. [TL;DR](#TL;DR)
|
| 56 |
+
1. [Model Details](#model-details)
|
| 57 |
+
2. [Usage](#usage)
|
| 58 |
+
3. [Uses](#uses)
|
| 59 |
+
4. [Bias, Risks, and Limitations](#bias-risks-and-limitations)
|
| 60 |
+
5. [Training Details](#training-details)
|
| 61 |
+
6. [Evaluation](#evaluation)
|
| 62 |
+
7. [Environmental Impact](#environmental-impact)
|
| 63 |
+
8. [Citation](#citation)
|
| 64 |
+
9. [Model Card Authors](#model-card-authors)
|
| 65 |
+
|
| 66 |
+
# TL;DR
|
| 67 |
+
|
| 68 |
+
If you already know T5, FLAN-T5 is just better at everything. For the same number of parameters, these models have been fine-tuned on more than 1000 additional tasks covering also more languages.
|
| 69 |
+
As mentioned in the first few lines of the abstract :
|
| 70 |
+
> Flan-PaLM 540B achieves state-of-the-art performance on several benchmarks, such as 75.2% on five-shot MMLU. We also publicly release Flan-T5 checkpoints,1 which achieve strong few-shot performance even compared to much larger models, such as PaLM 62B. Overall, instruction finetuning is a general method for improving the performance and usability of pretrained language models.
|
| 71 |
+
|
| 72 |
+
**Disclaimer**: Content from **this** model card has been written by the Hugging Face team, and parts of it were copy pasted from the [T5 model card](https://huggingface.co/t5-large).
|
| 73 |
+
|
| 74 |
+
# Model Details
|
| 75 |
+
|
| 76 |
+
## Model Description
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
- **Model type:** Language model
|
| 80 |
+
- **Language(s) (NLP):** English, Spanish, Japanese, Persian, Hindi, French, Chinese, Bengali, Gujarati, German, Telugu, Italian, Arabic, Polish, Tamil, Marathi, Malayalam, Oriya, Panjabi, Portuguese, Urdu, Galician, Hebrew, Korean, Catalan, Thai, Dutch, Indonesian, Vietnamese, Bulgarian, Filipino, Central Khmer, Lao, Turkish, Russian, Croatian, Swedish, Yoruba, Kurdish, Burmese, Malay, Czech, Finnish, Somali, Tagalog, Swahili, Sinhala, Kannada, Zhuang, Igbo, Xhosa, Romanian, Haitian, Estonian, Slovak, Lithuanian, Greek, Nepali, Assamese, Norwegian
|
| 81 |
+
- **License:** Apache 2.0
|
| 82 |
+
- **Related Models:** [All FLAN-T5 Checkpoints](https://huggingface.co/models?search=flan-t5)
|
| 83 |
+
- **Original Checkpoints:** [All Original FLAN-T5 Checkpoints](https://github.com/google-research/t5x/blob/main/docs/models.md#flan-t5-checkpoints)
|
| 84 |
+
- **Resources for more information:**
|
| 85 |
+
- [Research paper](https://arxiv.org/pdf/2210.11416.pdf)
|
| 86 |
+
- [GitHub Repo](https://github.com/google-research/t5x)
|
| 87 |
+
- [Hugging Face FLAN-T5 Docs (Similar to T5) ](https://huggingface.co/docs/transformers/model_doc/t5)
|
| 88 |
+
|
| 89 |
+
# Usage
|
| 90 |
+
|
| 91 |
+
Find below some example scripts on how to use the model in `transformers`:
|
| 92 |
+
|
| 93 |
+
## Using the Pytorch model
|
| 94 |
+
|
| 95 |
+
### Running the model on a CPU
|
| 96 |
+
|
| 97 |
+
<details>
|
| 98 |
+
<summary> Click to expand </summary>
|
| 99 |
+
|
| 100 |
+
```python
|
| 101 |
+
|
| 102 |
+
from transformers import T5Tokenizer, T5ForConditionalGeneration
|
| 103 |
+
|
| 104 |
+
tokenizer = T5Tokenizer.from_pretrained("google/flan-t5-large")
|
| 105 |
+
model = T5ForConditionalGeneration.from_pretrained("google/flan-t5-large")
|
| 106 |
+
|
| 107 |
+
input_text = "translate English to German: How old are you?"
|
| 108 |
+
input_ids = tokenizer(input_text, return_tensors="pt").input_ids
|
| 109 |
+
|
| 110 |
+
outputs = model.generate(input_ids)
|
| 111 |
+
print(tokenizer.decode(outputs[0]))
|
| 112 |
+
```
|
| 113 |
+
|
| 114 |
+
</details>
|
| 115 |
+
|
| 116 |
+
### Running the model on a GPU
|
| 117 |
+
|
| 118 |
+
<details>
|
| 119 |
+
<summary> Click to expand </summary>
|
| 120 |
+
|
| 121 |
+
```python
|
| 122 |
+
# pip install accelerate
|
| 123 |
+
from transformers import T5Tokenizer, T5ForConditionalGeneration
|
| 124 |
+
|
| 125 |
+
tokenizer = T5Tokenizer.from_pretrained("google/flan-t5-large")
|
| 126 |
+
model = T5ForConditionalGeneration.from_pretrained("google/flan-t5-large", device_map="auto")
|
| 127 |
+
|
| 128 |
+
input_text = "translate English to German: How old are you?"
|
| 129 |
+
input_ids = tokenizer(input_text, return_tensors="pt").input_ids.to("cuda")
|
| 130 |
+
|
| 131 |
+
outputs = model.generate(input_ids)
|
| 132 |
+
print(tokenizer.decode(outputs[0]))
|
| 133 |
+
```
|
| 134 |
+
|
| 135 |
+
</details>
|
| 136 |
+
|
| 137 |
+
### Running the model on a GPU using different precisions
|
| 138 |
+
|
| 139 |
+
#### FP16
|
| 140 |
+
|
| 141 |
+
<details>
|
| 142 |
+
<summary> Click to expand </summary>
|
| 143 |
+
|
| 144 |
+
```python
|
| 145 |
+
# pip install accelerate
|
| 146 |
+
import torch
|
| 147 |
+
from transformers import T5Tokenizer, T5ForConditionalGeneration
|
| 148 |
+
|
| 149 |
+
tokenizer = T5Tokenizer.from_pretrained("google/flan-t5-large")
|
| 150 |
+
model = T5ForConditionalGeneration.from_pretrained("google/flan-t5-large", device_map="auto", torch_dtype=torch.float16)
|
| 151 |
+
|
| 152 |
+
input_text = "translate English to German: How old are you?"
|
| 153 |
+
input_ids = tokenizer(input_text, return_tensors="pt").input_ids.to("cuda")
|
| 154 |
+
|
| 155 |
+
outputs = model.generate(input_ids)
|
| 156 |
+
print(tokenizer.decode(outputs[0]))
|
| 157 |
+
```
|
| 158 |
+
|
| 159 |
+
</details>
|
| 160 |
+
|
| 161 |
+
#### INT8
|
| 162 |
+
|
| 163 |
+
<details>
|
| 164 |
+
<summary> Click to expand </summary>
|
| 165 |
+
|
| 166 |
+
```python
|
| 167 |
+
# pip install bitsandbytes accelerate
|
| 168 |
+
from transformers import T5Tokenizer, T5ForConditionalGeneration
|
| 169 |
+
|
| 170 |
+
tokenizer = T5Tokenizer.from_pretrained("google/flan-t5-large")
|
| 171 |
+
model = T5ForConditionalGeneration.from_pretrained("google/flan-t5-large", device_map="auto", load_in_8bit=True)
|
| 172 |
+
|
| 173 |
+
input_text = "translate English to German: How old are you?"
|
| 174 |
+
input_ids = tokenizer(input_text, return_tensors="pt").input_ids.to("cuda")
|
| 175 |
+
|
| 176 |
+
outputs = model.generate(input_ids)
|
| 177 |
+
print(tokenizer.decode(outputs[0]))
|
| 178 |
+
```
|
| 179 |
+
|
| 180 |
+
</details>
|
| 181 |
+
|
| 182 |
+
# Uses
|
| 183 |
+
|
| 184 |
+
## Direct Use and Downstream Use
|
| 185 |
+
|
| 186 |
+
The authors write in [the original paper's model card](https://arxiv.org/pdf/2210.11416.pdf) that:
|
| 187 |
+
|
| 188 |
+
> The primary use is research on language models, including: research on zero-shot NLP tasks and in-context few-shot learning NLP tasks, such as reasoning, and question answering; advancing fairness and safety research, and understanding limitations of current large language models
|
| 189 |
+
|
| 190 |
+
See the [research paper](https://arxiv.org/pdf/2210.11416.pdf) for further details.
|
| 191 |
+
|
| 192 |
+
## Out-of-Scope Use
|
| 193 |
+
|
| 194 |
+
More information needed.
|
| 195 |
+
|
| 196 |
+
# Bias, Risks, and Limitations
|
| 197 |
+
|
| 198 |
+
The information below in this section are copied from the model's [official model card](https://arxiv.org/pdf/2210.11416.pdf):
|
| 199 |
+
|
| 200 |
+
> Language models, including Flan-T5, can potentially be used for language generation in a harmful way, according to Rae et al. (2021). Flan-T5 should not be used directly in any application, without a prior assessment of safety and fairness concerns specific to the application.
|
| 201 |
+
|
| 202 |
+
## Ethical considerations and risks
|
| 203 |
+
|
| 204 |
+
> Flan-T5 is fine-tuned on a large corpus of text data that was not filtered for explicit content or assessed for existing biases. As a result the model itself is potentially vulnerable to generating equivalently inappropriate content or replicating inherent biases in the underlying data.
|
| 205 |
+
|
| 206 |
+
## Known Limitations
|
| 207 |
+
|
| 208 |
+
> Flan-T5 has not been tested in real world applications.
|
| 209 |
+
|
| 210 |
+
## Sensitive Use:
|
| 211 |
+
|
| 212 |
+
> Flan-T5 should not be applied for any unacceptable use cases, e.g., generation of abusive speech.
|
| 213 |
+
|
| 214 |
+
# Training Details
|
| 215 |
+
|
| 216 |
+
## Training Data
|
| 217 |
+
|
| 218 |
+
The model was trained on a mixture of tasks, that includes the tasks described in the table below (from the original paper, figure 2):
|
| 219 |
+
|
| 220 |
+

|
| 221 |
+
|
| 222 |
+
|
| 223 |
+
## Training Procedure
|
| 224 |
+
|
| 225 |
+
According to the model card from the [original paper](https://arxiv.org/pdf/2210.11416.pdf):
|
| 226 |
+
|
| 227 |
+
> These models are based on pretrained T5 (Raffel et al., 2020) and fine-tuned with instructions for better zero-shot and few-shot performance. There is one fine-tuned Flan model per T5 model size.
|
| 228 |
+
|
| 229 |
+
The model has been trained on TPU v3 or TPU v4 pods, using [`t5x`](https://github.com/google-research/t5x) codebase together with [`jax`](https://github.com/google/jax).
|
| 230 |
+
|
| 231 |
+
|
| 232 |
+
# Evaluation
|
| 233 |
+
|
| 234 |
+
## Testing Data, Factors & Metrics
|
| 235 |
+
|
| 236 |
+
The authors evaluated the model on various tasks covering several languages (1836 in total). See the table below for some quantitative evaluation:
|
| 237 |
+

|
| 238 |
+
For full details, please check the [research paper](https://arxiv.org/pdf/2210.11416.pdf).
|
| 239 |
+
|
| 240 |
+
## Results
|
| 241 |
+
|
| 242 |
+
For full results for FLAN-T5-Large, see the [research paper](https://arxiv.org/pdf/2210.11416.pdf), Table 3.
|
| 243 |
+
|
| 244 |
+
# Environmental Impact
|
| 245 |
+
|
| 246 |
+
Carbon emissions can be estimated using the [Machine Learning Impact calculator](https://mlco2.github.io/impact#compute) presented in [Lacoste et al. (2019)](https://arxiv.org/abs/1910.09700).
|
| 247 |
+
|
| 248 |
+
- **Hardware Type:** Google Cloud TPU Pods - TPU v3 or TPU v4 | Number of chips β₯ 4.
|
| 249 |
+
- **Hours used:** More information needed
|
| 250 |
+
- **Cloud Provider:** GCP
|
| 251 |
+
- **Compute Region:** More information needed
|
| 252 |
+
- **Carbon Emitted:** More information needed
|
| 253 |
+
|
| 254 |
+
# Citation
|
| 255 |
+
|
| 256 |
+
**BibTeX:**
|
| 257 |
+
|
| 258 |
+
```bibtex
|
| 259 |
+
@misc{https://doi.org/10.48550/arxiv.2210.11416,
|
| 260 |
+
doi = {10.48550/ARXIV.2210.11416},
|
| 261 |
+
|
| 262 |
+
url = {https://arxiv.org/abs/2210.11416},
|
| 263 |
+
|
| 264 |
+
author = {Chung, Hyung Won and Hou, Le and Longpre, Shayne and Zoph, Barret and Tay, Yi and Fedus, William and Li, Eric and Wang, Xuezhi and Dehghani, Mostafa and Brahma, Siddhartha and Webson, Albert and Gu, Shixiang Shane and Dai, Zhuyun and Suzgun, Mirac and Chen, Xinyun and Chowdhery, Aakanksha and Narang, Sharan and Mishra, Gaurav and Yu, Adams and Zhao, Vincent and Huang, Yanping and Dai, Andrew and Yu, Hongkun and Petrov, Slav and Chi, Ed H. and Dean, Jeff and Devlin, Jacob and Roberts, Adam and Zhou, Denny and Le, Quoc V. and Wei, Jason},
|
| 265 |
+
|
| 266 |
+
keywords = {Machine Learning (cs.LG), Computation and Language (cs.CL), FOS: Computer and information sciences, FOS: Computer and information sciences},
|
| 267 |
+
|
| 268 |
+
title = {Scaling Instruction-Finetuned Language Models},
|
| 269 |
+
|
| 270 |
+
publisher = {arXiv},
|
| 271 |
+
|
| 272 |
+
year = {2022},
|
| 273 |
+
|
| 274 |
+
copyright = {Creative Commons Attribution 4.0 International}
|
| 275 |
+
}
|
| 276 |
+
```
|
Looped-DiT-B-16/text_encoder/config.json
ADDED
|
@@ -0,0 +1,28 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"architectures": [
|
| 3 |
+
"T5ForConditionalGeneration"
|
| 4 |
+
],
|
| 5 |
+
"d_ff": 2816,
|
| 6 |
+
"d_kv": 64,
|
| 7 |
+
"d_model": 1024,
|
| 8 |
+
"decoder_start_token_id": 0,
|
| 9 |
+
"dropout_rate": 0.1,
|
| 10 |
+
"eos_token_id": 1,
|
| 11 |
+
"feed_forward_proj": "gated-gelu",
|
| 12 |
+
"initializer_factor": 1.0,
|
| 13 |
+
"is_encoder_decoder": true,
|
| 14 |
+
"layer_norm_epsilon": 1e-06,
|
| 15 |
+
"model_type": "t5",
|
| 16 |
+
"n_positions": 512,
|
| 17 |
+
"num_decoder_layers": 24,
|
| 18 |
+
"num_heads": 16,
|
| 19 |
+
"num_layers": 24,
|
| 20 |
+
"output_past": true,
|
| 21 |
+
"pad_token_id": 0,
|
| 22 |
+
"relative_attention_max_distance": 128,
|
| 23 |
+
"relative_attention_num_buckets": 32,
|
| 24 |
+
"tie_word_embeddings": false,
|
| 25 |
+
"transformers_version": "4.23.1",
|
| 26 |
+
"use_cache": true,
|
| 27 |
+
"vocab_size": 32128
|
| 28 |
+
}
|
Looped-DiT-B-16/text_encoder/generation_config.json
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"_from_model_config": true,
|
| 3 |
+
"decoder_start_token_id": 0,
|
| 4 |
+
"eos_token_id": 1,
|
| 5 |
+
"pad_token_id": 0,
|
| 6 |
+
"transformers_version": "4.27.0.dev0"
|
| 7 |
+
}
|
Looped-DiT-B-16/text_encoder/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:149fe330e51e007c6ea9c6089f334ca82fcbe06b0df27e578e90934bcd327f73
|
| 3 |
+
size 3142856782
|
Looped-DiT-B-16/text_encoder/special_tokens_map.json
ADDED
|
@@ -0,0 +1,107 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"additional_special_tokens": [
|
| 3 |
+
"<extra_id_0>",
|
| 4 |
+
"<extra_id_1>",
|
| 5 |
+
"<extra_id_2>",
|
| 6 |
+
"<extra_id_3>",
|
| 7 |
+
"<extra_id_4>",
|
| 8 |
+
"<extra_id_5>",
|
| 9 |
+
"<extra_id_6>",
|
| 10 |
+
"<extra_id_7>",
|
| 11 |
+
"<extra_id_8>",
|
| 12 |
+
"<extra_id_9>",
|
| 13 |
+
"<extra_id_10>",
|
| 14 |
+
"<extra_id_11>",
|
| 15 |
+
"<extra_id_12>",
|
| 16 |
+
"<extra_id_13>",
|
| 17 |
+
"<extra_id_14>",
|
| 18 |
+
"<extra_id_15>",
|
| 19 |
+
"<extra_id_16>",
|
| 20 |
+
"<extra_id_17>",
|
| 21 |
+
"<extra_id_18>",
|
| 22 |
+
"<extra_id_19>",
|
| 23 |
+
"<extra_id_20>",
|
| 24 |
+
"<extra_id_21>",
|
| 25 |
+
"<extra_id_22>",
|
| 26 |
+
"<extra_id_23>",
|
| 27 |
+
"<extra_id_24>",
|
| 28 |
+
"<extra_id_25>",
|
| 29 |
+
"<extra_id_26>",
|
| 30 |
+
"<extra_id_27>",
|
| 31 |
+
"<extra_id_28>",
|
| 32 |
+
"<extra_id_29>",
|
| 33 |
+
"<extra_id_30>",
|
| 34 |
+
"<extra_id_31>",
|
| 35 |
+
"<extra_id_32>",
|
| 36 |
+
"<extra_id_33>",
|
| 37 |
+
"<extra_id_34>",
|
| 38 |
+
"<extra_id_35>",
|
| 39 |
+
"<extra_id_36>",
|
| 40 |
+
"<extra_id_37>",
|
| 41 |
+
"<extra_id_38>",
|
| 42 |
+
"<extra_id_39>",
|
| 43 |
+
"<extra_id_40>",
|
| 44 |
+
"<extra_id_41>",
|
| 45 |
+
"<extra_id_42>",
|
| 46 |
+
"<extra_id_43>",
|
| 47 |
+
"<extra_id_44>",
|
| 48 |
+
"<extra_id_45>",
|
| 49 |
+
"<extra_id_46>",
|
| 50 |
+
"<extra_id_47>",
|
| 51 |
+
"<extra_id_48>",
|
| 52 |
+
"<extra_id_49>",
|
| 53 |
+
"<extra_id_50>",
|
| 54 |
+
"<extra_id_51>",
|
| 55 |
+
"<extra_id_52>",
|
| 56 |
+
"<extra_id_53>",
|
| 57 |
+
"<extra_id_54>",
|
| 58 |
+
"<extra_id_55>",
|
| 59 |
+
"<extra_id_56>",
|
| 60 |
+
"<extra_id_57>",
|
| 61 |
+
"<extra_id_58>",
|
| 62 |
+
"<extra_id_59>",
|
| 63 |
+
"<extra_id_60>",
|
| 64 |
+
"<extra_id_61>",
|
| 65 |
+
"<extra_id_62>",
|
| 66 |
+
"<extra_id_63>",
|
| 67 |
+
"<extra_id_64>",
|
| 68 |
+
"<extra_id_65>",
|
| 69 |
+
"<extra_id_66>",
|
| 70 |
+
"<extra_id_67>",
|
| 71 |
+
"<extra_id_68>",
|
| 72 |
+
"<extra_id_69>",
|
| 73 |
+
"<extra_id_70>",
|
| 74 |
+
"<extra_id_71>",
|
| 75 |
+
"<extra_id_72>",
|
| 76 |
+
"<extra_id_73>",
|
| 77 |
+
"<extra_id_74>",
|
| 78 |
+
"<extra_id_75>",
|
| 79 |
+
"<extra_id_76>",
|
| 80 |
+
"<extra_id_77>",
|
| 81 |
+
"<extra_id_78>",
|
| 82 |
+
"<extra_id_79>",
|
| 83 |
+
"<extra_id_80>",
|
| 84 |
+
"<extra_id_81>",
|
| 85 |
+
"<extra_id_82>",
|
| 86 |
+
"<extra_id_83>",
|
| 87 |
+
"<extra_id_84>",
|
| 88 |
+
"<extra_id_85>",
|
| 89 |
+
"<extra_id_86>",
|
| 90 |
+
"<extra_id_87>",
|
| 91 |
+
"<extra_id_88>",
|
| 92 |
+
"<extra_id_89>",
|
| 93 |
+
"<extra_id_90>",
|
| 94 |
+
"<extra_id_91>",
|
| 95 |
+
"<extra_id_92>",
|
| 96 |
+
"<extra_id_93>",
|
| 97 |
+
"<extra_id_94>",
|
| 98 |
+
"<extra_id_95>",
|
| 99 |
+
"<extra_id_96>",
|
| 100 |
+
"<extra_id_97>",
|
| 101 |
+
"<extra_id_98>",
|
| 102 |
+
"<extra_id_99>"
|
| 103 |
+
],
|
| 104 |
+
"eos_token": "</s>",
|
| 105 |
+
"pad_token": "<pad>",
|
| 106 |
+
"unk_token": "<unk>"
|
| 107 |
+
}
|
Looped-DiT-B-16/text_encoder/spiece.model
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:89fa65b45c6c46d9ff3ecaf7a4eeff28d758a92d87bda5102dcb8141a0c051d3
|
| 3 |
+
size 859107
|
Looped-DiT-B-16/text_encoder/tokenizer.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
Looped-DiT-B-16/text_encoder/tokenizer_config.json
ADDED
|
@@ -0,0 +1,113 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"additional_special_tokens": [
|
| 3 |
+
"<extra_id_0>",
|
| 4 |
+
"<extra_id_1>",
|
| 5 |
+
"<extra_id_2>",
|
| 6 |
+
"<extra_id_3>",
|
| 7 |
+
"<extra_id_4>",
|
| 8 |
+
"<extra_id_5>",
|
| 9 |
+
"<extra_id_6>",
|
| 10 |
+
"<extra_id_7>",
|
| 11 |
+
"<extra_id_8>",
|
| 12 |
+
"<extra_id_9>",
|
| 13 |
+
"<extra_id_10>",
|
| 14 |
+
"<extra_id_11>",
|
| 15 |
+
"<extra_id_12>",
|
| 16 |
+
"<extra_id_13>",
|
| 17 |
+
"<extra_id_14>",
|
| 18 |
+
"<extra_id_15>",
|
| 19 |
+
"<extra_id_16>",
|
| 20 |
+
"<extra_id_17>",
|
| 21 |
+
"<extra_id_18>",
|
| 22 |
+
"<extra_id_19>",
|
| 23 |
+
"<extra_id_20>",
|
| 24 |
+
"<extra_id_21>",
|
| 25 |
+
"<extra_id_22>",
|
| 26 |
+
"<extra_id_23>",
|
| 27 |
+
"<extra_id_24>",
|
| 28 |
+
"<extra_id_25>",
|
| 29 |
+
"<extra_id_26>",
|
| 30 |
+
"<extra_id_27>",
|
| 31 |
+
"<extra_id_28>",
|
| 32 |
+
"<extra_id_29>",
|
| 33 |
+
"<extra_id_30>",
|
| 34 |
+
"<extra_id_31>",
|
| 35 |
+
"<extra_id_32>",
|
| 36 |
+
"<extra_id_33>",
|
| 37 |
+
"<extra_id_34>",
|
| 38 |
+
"<extra_id_35>",
|
| 39 |
+
"<extra_id_36>",
|
| 40 |
+
"<extra_id_37>",
|
| 41 |
+
"<extra_id_38>",
|
| 42 |
+
"<extra_id_39>",
|
| 43 |
+
"<extra_id_40>",
|
| 44 |
+
"<extra_id_41>",
|
| 45 |
+
"<extra_id_42>",
|
| 46 |
+
"<extra_id_43>",
|
| 47 |
+
"<extra_id_44>",
|
| 48 |
+
"<extra_id_45>",
|
| 49 |
+
"<extra_id_46>",
|
| 50 |
+
"<extra_id_47>",
|
| 51 |
+
"<extra_id_48>",
|
| 52 |
+
"<extra_id_49>",
|
| 53 |
+
"<extra_id_50>",
|
| 54 |
+
"<extra_id_51>",
|
| 55 |
+
"<extra_id_52>",
|
| 56 |
+
"<extra_id_53>",
|
| 57 |
+
"<extra_id_54>",
|
| 58 |
+
"<extra_id_55>",
|
| 59 |
+
"<extra_id_56>",
|
| 60 |
+
"<extra_id_57>",
|
| 61 |
+
"<extra_id_58>",
|
| 62 |
+
"<extra_id_59>",
|
| 63 |
+
"<extra_id_60>",
|
| 64 |
+
"<extra_id_61>",
|
| 65 |
+
"<extra_id_62>",
|
| 66 |
+
"<extra_id_63>",
|
| 67 |
+
"<extra_id_64>",
|
| 68 |
+
"<extra_id_65>",
|
| 69 |
+
"<extra_id_66>",
|
| 70 |
+
"<extra_id_67>",
|
| 71 |
+
"<extra_id_68>",
|
| 72 |
+
"<extra_id_69>",
|
| 73 |
+
"<extra_id_70>",
|
| 74 |
+
"<extra_id_71>",
|
| 75 |
+
"<extra_id_72>",
|
| 76 |
+
"<extra_id_73>",
|
| 77 |
+
"<extra_id_74>",
|
| 78 |
+
"<extra_id_75>",
|
| 79 |
+
"<extra_id_76>",
|
| 80 |
+
"<extra_id_77>",
|
| 81 |
+
"<extra_id_78>",
|
| 82 |
+
"<extra_id_79>",
|
| 83 |
+
"<extra_id_80>",
|
| 84 |
+
"<extra_id_81>",
|
| 85 |
+
"<extra_id_82>",
|
| 86 |
+
"<extra_id_83>",
|
| 87 |
+
"<extra_id_84>",
|
| 88 |
+
"<extra_id_85>",
|
| 89 |
+
"<extra_id_86>",
|
| 90 |
+
"<extra_id_87>",
|
| 91 |
+
"<extra_id_88>",
|
| 92 |
+
"<extra_id_89>",
|
| 93 |
+
"<extra_id_90>",
|
| 94 |
+
"<extra_id_91>",
|
| 95 |
+
"<extra_id_92>",
|
| 96 |
+
"<extra_id_93>",
|
| 97 |
+
"<extra_id_94>",
|
| 98 |
+
"<extra_id_95>",
|
| 99 |
+
"<extra_id_96>",
|
| 100 |
+
"<extra_id_97>",
|
| 101 |
+
"<extra_id_98>",
|
| 102 |
+
"<extra_id_99>"
|
| 103 |
+
],
|
| 104 |
+
"eos_token": "</s>",
|
| 105 |
+
"extra_ids": 100,
|
| 106 |
+
"model_max_length": 256,
|
| 107 |
+
"name_or_path": "google/t5-v1_1-large",
|
| 108 |
+
"pad_token": "<pad>",
|
| 109 |
+
"sp_model_kwargs": {},
|
| 110 |
+
"special_tokens_map_file": "/home/younes_huggingface_co/.cache/huggingface/hub/models--google--t5-v1_1-large/snapshots/314bc112b191ec17b625ba81438dc73d6c23659d/special_tokens_map.json",
|
| 111 |
+
"tokenizer_class": "T5Tokenizer",
|
| 112 |
+
"unk_token": "<unk>"
|
| 113 |
+
}
|
Looped-DiT-B-16/tokenizer/special_tokens_map.json
ADDED
|
@@ -0,0 +1,107 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"additional_special_tokens": [
|
| 3 |
+
"<extra_id_0>",
|
| 4 |
+
"<extra_id_1>",
|
| 5 |
+
"<extra_id_2>",
|
| 6 |
+
"<extra_id_3>",
|
| 7 |
+
"<extra_id_4>",
|
| 8 |
+
"<extra_id_5>",
|
| 9 |
+
"<extra_id_6>",
|
| 10 |
+
"<extra_id_7>",
|
| 11 |
+
"<extra_id_8>",
|
| 12 |
+
"<extra_id_9>",
|
| 13 |
+
"<extra_id_10>",
|
| 14 |
+
"<extra_id_11>",
|
| 15 |
+
"<extra_id_12>",
|
| 16 |
+
"<extra_id_13>",
|
| 17 |
+
"<extra_id_14>",
|
| 18 |
+
"<extra_id_15>",
|
| 19 |
+
"<extra_id_16>",
|
| 20 |
+
"<extra_id_17>",
|
| 21 |
+
"<extra_id_18>",
|
| 22 |
+
"<extra_id_19>",
|
| 23 |
+
"<extra_id_20>",
|
| 24 |
+
"<extra_id_21>",
|
| 25 |
+
"<extra_id_22>",
|
| 26 |
+
"<extra_id_23>",
|
| 27 |
+
"<extra_id_24>",
|
| 28 |
+
"<extra_id_25>",
|
| 29 |
+
"<extra_id_26>",
|
| 30 |
+
"<extra_id_27>",
|
| 31 |
+
"<extra_id_28>",
|
| 32 |
+
"<extra_id_29>",
|
| 33 |
+
"<extra_id_30>",
|
| 34 |
+
"<extra_id_31>",
|
| 35 |
+
"<extra_id_32>",
|
| 36 |
+
"<extra_id_33>",
|
| 37 |
+
"<extra_id_34>",
|
| 38 |
+
"<extra_id_35>",
|
| 39 |
+
"<extra_id_36>",
|
| 40 |
+
"<extra_id_37>",
|
| 41 |
+
"<extra_id_38>",
|
| 42 |
+
"<extra_id_39>",
|
| 43 |
+
"<extra_id_40>",
|
| 44 |
+
"<extra_id_41>",
|
| 45 |
+
"<extra_id_42>",
|
| 46 |
+
"<extra_id_43>",
|
| 47 |
+
"<extra_id_44>",
|
| 48 |
+
"<extra_id_45>",
|
| 49 |
+
"<extra_id_46>",
|
| 50 |
+
"<extra_id_47>",
|
| 51 |
+
"<extra_id_48>",
|
| 52 |
+
"<extra_id_49>",
|
| 53 |
+
"<extra_id_50>",
|
| 54 |
+
"<extra_id_51>",
|
| 55 |
+
"<extra_id_52>",
|
| 56 |
+
"<extra_id_53>",
|
| 57 |
+
"<extra_id_54>",
|
| 58 |
+
"<extra_id_55>",
|
| 59 |
+
"<extra_id_56>",
|
| 60 |
+
"<extra_id_57>",
|
| 61 |
+
"<extra_id_58>",
|
| 62 |
+
"<extra_id_59>",
|
| 63 |
+
"<extra_id_60>",
|
| 64 |
+
"<extra_id_61>",
|
| 65 |
+
"<extra_id_62>",
|
| 66 |
+
"<extra_id_63>",
|
| 67 |
+
"<extra_id_64>",
|
| 68 |
+
"<extra_id_65>",
|
| 69 |
+
"<extra_id_66>",
|
| 70 |
+
"<extra_id_67>",
|
| 71 |
+
"<extra_id_68>",
|
| 72 |
+
"<extra_id_69>",
|
| 73 |
+
"<extra_id_70>",
|
| 74 |
+
"<extra_id_71>",
|
| 75 |
+
"<extra_id_72>",
|
| 76 |
+
"<extra_id_73>",
|
| 77 |
+
"<extra_id_74>",
|
| 78 |
+
"<extra_id_75>",
|
| 79 |
+
"<extra_id_76>",
|
| 80 |
+
"<extra_id_77>",
|
| 81 |
+
"<extra_id_78>",
|
| 82 |
+
"<extra_id_79>",
|
| 83 |
+
"<extra_id_80>",
|
| 84 |
+
"<extra_id_81>",
|
| 85 |
+
"<extra_id_82>",
|
| 86 |
+
"<extra_id_83>",
|
| 87 |
+
"<extra_id_84>",
|
| 88 |
+
"<extra_id_85>",
|
| 89 |
+
"<extra_id_86>",
|
| 90 |
+
"<extra_id_87>",
|
| 91 |
+
"<extra_id_88>",
|
| 92 |
+
"<extra_id_89>",
|
| 93 |
+
"<extra_id_90>",
|
| 94 |
+
"<extra_id_91>",
|
| 95 |
+
"<extra_id_92>",
|
| 96 |
+
"<extra_id_93>",
|
| 97 |
+
"<extra_id_94>",
|
| 98 |
+
"<extra_id_95>",
|
| 99 |
+
"<extra_id_96>",
|
| 100 |
+
"<extra_id_97>",
|
| 101 |
+
"<extra_id_98>",
|
| 102 |
+
"<extra_id_99>"
|
| 103 |
+
],
|
| 104 |
+
"eos_token": "</s>",
|
| 105 |
+
"pad_token": "<pad>",
|
| 106 |
+
"unk_token": "<unk>"
|
| 107 |
+
}
|
Looped-DiT-B-16/tokenizer/spiece.model
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:89fa65b45c6c46d9ff3ecaf7a4eeff28d758a92d87bda5102dcb8141a0c051d3
|
| 3 |
+
size 859107
|
Looped-DiT-B-16/tokenizer/tokenizer.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
Looped-DiT-B-16/tokenizer/tokenizer_config.json
ADDED
|
@@ -0,0 +1,113 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"additional_special_tokens": [
|
| 3 |
+
"<extra_id_0>",
|
| 4 |
+
"<extra_id_1>",
|
| 5 |
+
"<extra_id_2>",
|
| 6 |
+
"<extra_id_3>",
|
| 7 |
+
"<extra_id_4>",
|
| 8 |
+
"<extra_id_5>",
|
| 9 |
+
"<extra_id_6>",
|
| 10 |
+
"<extra_id_7>",
|
| 11 |
+
"<extra_id_8>",
|
| 12 |
+
"<extra_id_9>",
|
| 13 |
+
"<extra_id_10>",
|
| 14 |
+
"<extra_id_11>",
|
| 15 |
+
"<extra_id_12>",
|
| 16 |
+
"<extra_id_13>",
|
| 17 |
+
"<extra_id_14>",
|
| 18 |
+
"<extra_id_15>",
|
| 19 |
+
"<extra_id_16>",
|
| 20 |
+
"<extra_id_17>",
|
| 21 |
+
"<extra_id_18>",
|
| 22 |
+
"<extra_id_19>",
|
| 23 |
+
"<extra_id_20>",
|
| 24 |
+
"<extra_id_21>",
|
| 25 |
+
"<extra_id_22>",
|
| 26 |
+
"<extra_id_23>",
|
| 27 |
+
"<extra_id_24>",
|
| 28 |
+
"<extra_id_25>",
|
| 29 |
+
"<extra_id_26>",
|
| 30 |
+
"<extra_id_27>",
|
| 31 |
+
"<extra_id_28>",
|
| 32 |
+
"<extra_id_29>",
|
| 33 |
+
"<extra_id_30>",
|
| 34 |
+
"<extra_id_31>",
|
| 35 |
+
"<extra_id_32>",
|
| 36 |
+
"<extra_id_33>",
|
| 37 |
+
"<extra_id_34>",
|
| 38 |
+
"<extra_id_35>",
|
| 39 |
+
"<extra_id_36>",
|
| 40 |
+
"<extra_id_37>",
|
| 41 |
+
"<extra_id_38>",
|
| 42 |
+
"<extra_id_39>",
|
| 43 |
+
"<extra_id_40>",
|
| 44 |
+
"<extra_id_41>",
|
| 45 |
+
"<extra_id_42>",
|
| 46 |
+
"<extra_id_43>",
|
| 47 |
+
"<extra_id_44>",
|
| 48 |
+
"<extra_id_45>",
|
| 49 |
+
"<extra_id_46>",
|
| 50 |
+
"<extra_id_47>",
|
| 51 |
+
"<extra_id_48>",
|
| 52 |
+
"<extra_id_49>",
|
| 53 |
+
"<extra_id_50>",
|
| 54 |
+
"<extra_id_51>",
|
| 55 |
+
"<extra_id_52>",
|
| 56 |
+
"<extra_id_53>",
|
| 57 |
+
"<extra_id_54>",
|
| 58 |
+
"<extra_id_55>",
|
| 59 |
+
"<extra_id_56>",
|
| 60 |
+
"<extra_id_57>",
|
| 61 |
+
"<extra_id_58>",
|
| 62 |
+
"<extra_id_59>",
|
| 63 |
+
"<extra_id_60>",
|
| 64 |
+
"<extra_id_61>",
|
| 65 |
+
"<extra_id_62>",
|
| 66 |
+
"<extra_id_63>",
|
| 67 |
+
"<extra_id_64>",
|
| 68 |
+
"<extra_id_65>",
|
| 69 |
+
"<extra_id_66>",
|
| 70 |
+
"<extra_id_67>",
|
| 71 |
+
"<extra_id_68>",
|
| 72 |
+
"<extra_id_69>",
|
| 73 |
+
"<extra_id_70>",
|
| 74 |
+
"<extra_id_71>",
|
| 75 |
+
"<extra_id_72>",
|
| 76 |
+
"<extra_id_73>",
|
| 77 |
+
"<extra_id_74>",
|
| 78 |
+
"<extra_id_75>",
|
| 79 |
+
"<extra_id_76>",
|
| 80 |
+
"<extra_id_77>",
|
| 81 |
+
"<extra_id_78>",
|
| 82 |
+
"<extra_id_79>",
|
| 83 |
+
"<extra_id_80>",
|
| 84 |
+
"<extra_id_81>",
|
| 85 |
+
"<extra_id_82>",
|
| 86 |
+
"<extra_id_83>",
|
| 87 |
+
"<extra_id_84>",
|
| 88 |
+
"<extra_id_85>",
|
| 89 |
+
"<extra_id_86>",
|
| 90 |
+
"<extra_id_87>",
|
| 91 |
+
"<extra_id_88>",
|
| 92 |
+
"<extra_id_89>",
|
| 93 |
+
"<extra_id_90>",
|
| 94 |
+
"<extra_id_91>",
|
| 95 |
+
"<extra_id_92>",
|
| 96 |
+
"<extra_id_93>",
|
| 97 |
+
"<extra_id_94>",
|
| 98 |
+
"<extra_id_95>",
|
| 99 |
+
"<extra_id_96>",
|
| 100 |
+
"<extra_id_97>",
|
| 101 |
+
"<extra_id_98>",
|
| 102 |
+
"<extra_id_99>"
|
| 103 |
+
],
|
| 104 |
+
"eos_token": "</s>",
|
| 105 |
+
"extra_ids": 100,
|
| 106 |
+
"model_max_length": 256,
|
| 107 |
+
"name_or_path": "google/t5-v1_1-large",
|
| 108 |
+
"pad_token": "<pad>",
|
| 109 |
+
"sp_model_kwargs": {},
|
| 110 |
+
"special_tokens_map_file": "/home/younes_huggingface_co/.cache/huggingface/hub/models--google--t5-v1_1-large/snapshots/314bc112b191ec17b625ba81438dc73d6c23659d/special_tokens_map.json",
|
| 111 |
+
"tokenizer_class": "T5Tokenizer",
|
| 112 |
+
"unk_token": "<unk>"
|
| 113 |
+
}
|
Looped-DiT-B-16/transformer/config.json
ADDED
|
@@ -0,0 +1,23 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"_class_name": "LoopedDiTTransformer2DModel",
|
| 3 |
+
"_diffusers_version": "0.39.0",
|
| 4 |
+
"head_dim": 64,
|
| 5 |
+
"hidden_size": 768,
|
| 6 |
+
"image_size": 512,
|
| 7 |
+
"in_channels": 3,
|
| 8 |
+
"loop_split": [
|
| 9 |
+
6,
|
| 10 |
+
5,
|
| 11 |
+
6
|
| 12 |
+
],
|
| 13 |
+
"mlp_ratio": 2.6667,
|
| 14 |
+
"num_heads": 12,
|
| 15 |
+
"num_loops": 4,
|
| 16 |
+
"patch_size": 16,
|
| 17 |
+
"pca_channels": 128,
|
| 18 |
+
"share_loop_weights": true,
|
| 19 |
+
"text_dim": 1024,
|
| 20 |
+
"text_preamble_depth": 2,
|
| 21 |
+
"use_attn_gate": false,
|
| 22 |
+
"use_xsa": true
|
| 23 |
+
}
|
Looped-DiT-B-16/transformer/diffusion_pytorch_model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:885152c62ff0538a91ee6fd6ca1699bebd7972468aed430446a58a63f530cabf
|
| 3 |
+
size 1035851844
|
Looped-DiT-B-16/transformer/transformer_looped_dit.py
ADDED
|
@@ -0,0 +1,417 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Looped-DiT denoiser as a diffusers model.
|
| 2 |
+
|
| 3 |
+
The MiniT2I denoiser, a pixel-space variant of MMDiT (patchified image tokens
|
| 4 |
+
and T5 text tokens with modality-specific weights and joint attention, no
|
| 5 |
+
timestep conditioning), whose double-stream blocks are split into three stages:
|
| 6 |
+
|
| 7 |
+
pre-loop A blocks[:pre] run once
|
| 8 |
+
looped B the next `core` blocks run N times
|
| 9 |
+
post-loop C the last `post` blocks run once, followed by the head
|
| 10 |
+
|
| 11 |
+
h_0 = A(x), h_r = B(h_{r-1}), x0_hat(r) = C(h_r), r = 1..N
|
| 12 |
+
|
| 13 |
+
With shared weights (the default) B is one set of blocks reused N times, so the
|
| 14 |
+
loop adds depth but no weights. Any loop state h_r can be decoded through C:
|
| 15 |
+
deep supervision trains those intermediate exits, and inference can run with a
|
| 16 |
+
different loop depth. XSA and the attention gate act on the looped blocks only.
|
| 17 |
+
"""
|
| 18 |
+
|
| 19 |
+
from __future__ import annotations
|
| 20 |
+
|
| 21 |
+
import math
|
| 22 |
+
|
| 23 |
+
import torch
|
| 24 |
+
import torch.nn.functional as F
|
| 25 |
+
from diffusers.configuration_utils import ConfigMixin
|
| 26 |
+
from diffusers.models.modeling_utils import ModelMixin
|
| 27 |
+
from torch import nn
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
class RMSNorm(nn.Module):
|
| 31 |
+
def __init__(self, dim: int, eps: float = 1e-6):
|
| 32 |
+
super().__init__()
|
| 33 |
+
self.eps = eps
|
| 34 |
+
self.weight = nn.Parameter(torch.ones(dim))
|
| 35 |
+
|
| 36 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 37 |
+
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps) * self.weight
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
class SwiGLU(nn.Module):
|
| 41 |
+
def __init__(self, dim: int, hidden_dim: int):
|
| 42 |
+
super().__init__()
|
| 43 |
+
hidden_dim = math.ceil(hidden_dim / 8) * 8
|
| 44 |
+
self.w1 = nn.Linear(dim, hidden_dim, bias=False)
|
| 45 |
+
self.w3 = nn.Linear(dim, hidden_dim, bias=False)
|
| 46 |
+
self.w2 = nn.Linear(hidden_dim, dim, bias=False)
|
| 47 |
+
for layer in (self.w1, self.w3, self.w2):
|
| 48 |
+
nn.init.xavier_uniform_(layer.weight)
|
| 49 |
+
|
| 50 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 51 |
+
return self.w2(F.silu(self.w1(x)) * self.w3(x))
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
# ---------------------------------------------------------------------------
|
| 55 |
+
# Rotary position embeddings: 1D over text positions, 2D over the patch grid.
|
| 56 |
+
# ---------------------------------------------------------------------------
|
| 57 |
+
|
| 58 |
+
_ROPE_CACHE: dict[tuple, tuple[torch.Tensor, torch.Tensor]] = {}
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
def _autocast_state() -> tuple:
|
| 62 |
+
try:
|
| 63 |
+
enabled = torch.is_autocast_enabled("cuda")
|
| 64 |
+
return enabled, torch.get_autocast_dtype("cuda") if enabled else None
|
| 65 |
+
except TypeError: # torch < 2.4
|
| 66 |
+
enabled = torch.is_autocast_enabled()
|
| 67 |
+
return enabled, torch.get_autocast_gpu_dtype() if enabled else None
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
def _rope_tables(grid: int | None, n: int, d: int, device, dtype, theta: float = 10000.0):
|
| 71 |
+
"""cos/sin tables for 1D (grid=None) or 2D rotary embeddings.
|
| 72 |
+
|
| 73 |
+
Tables are built under whatever autocast state the caller runs in (under bf16
|
| 74 |
+
autocast the angle products are computed in bf16, which is what the models
|
| 75 |
+
were trained with), so that state is part of the cache key.
|
| 76 |
+
"""
|
| 77 |
+
key = (grid, n, d, str(device), dtype, _autocast_state())
|
| 78 |
+
if key not in _ROPE_CACHE:
|
| 79 |
+
if grid is None:
|
| 80 |
+
inv = 1.0 / (theta ** (torch.arange(0, d, 2, device=device, dtype=torch.float32) / d))
|
| 81 |
+
pos = torch.arange(n, device=device, dtype=torch.float32)
|
| 82 |
+
angles = torch.einsum("n,f->nf", pos, inv)
|
| 83 |
+
angles = torch.cat([angles, angles], dim=-1)
|
| 84 |
+
else:
|
| 85 |
+
half = d // 2
|
| 86 |
+
inv = 1.0 / (theta ** (torch.arange(0, half, 2, device=device, dtype=torch.float32) / half))
|
| 87 |
+
freqs = torch.einsum("n,f->nf", torch.arange(grid, device=device, dtype=torch.float32), inv)
|
| 88 |
+
f_h, f_w = torch.broadcast_tensors(freqs[:, None, :], freqs[None, :, :])
|
| 89 |
+
angles = torch.cat([f_h, f_w], dim=-1)
|
| 90 |
+
angles = torch.cat([angles, angles], dim=-1).reshape(n, d)
|
| 91 |
+
_ROPE_CACHE[key] = (angles.cos()[None, None].to(dtype), angles.sin()[None, None].to(dtype))
|
| 92 |
+
return _ROPE_CACHE[key]
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
def rotate_half(x: torch.Tensor) -> torch.Tensor:
|
| 96 |
+
x1, x2 = x.chunk(2, dim=-1)
|
| 97 |
+
return torch.cat([-x2, x1], dim=-1)
|
| 98 |
+
|
| 99 |
+
|
| 100 |
+
def apply_rope(x: torch.Tensor, grid: int | None = None) -> torch.Tensor:
|
| 101 |
+
"""x: [batch, heads, tokens, head_dim]."""
|
| 102 |
+
cos, sin = _rope_tables(grid, x.shape[2], x.shape[3], x.device, x.dtype)
|
| 103 |
+
return x * cos + rotate_half(x) * sin
|
| 104 |
+
|
| 105 |
+
|
| 106 |
+
def sincos_2d(dim: int, grid: int) -> torch.Tensor:
|
| 107 |
+
y, x = torch.meshgrid(torch.arange(grid), torch.arange(grid), indexing="ij")
|
| 108 |
+
omega = 1.0 / (10000 ** (torch.arange(dim // 4, dtype=torch.float32) / (dim // 4)))
|
| 109 |
+
out_y = torch.einsum("n,d->nd", y.flatten().float(), omega)
|
| 110 |
+
out_x = torch.einsum("n,d->nd", x.flatten().float(), omega)
|
| 111 |
+
return torch.cat([out_x.sin(), out_x.cos(), out_y.sin(), out_y.cos()], dim=1)
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
def exclusive_self_attention(out: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
|
| 115 |
+
"""XSA (Zhai, 2026: https://arxiv.org/abs/2603.09078):
|
| 116 |
+
remove from each token's attention output the component along that token's
|
| 117 |
+
own value vector, so attention only writes content from other tokens.
|
| 118 |
+
`out` and `v` are token-aligned, heads first."""
|
| 119 |
+
v_hat = F.normalize(v.float(), dim=-1)
|
| 120 |
+
out_f = out.float()
|
| 121 |
+
return (out_f - (out_f * v_hat).sum(dim=-1, keepdim=True) * v_hat).to(out.dtype)
|
| 122 |
+
|
| 123 |
+
|
| 124 |
+
# ---------------------------------------------------------------------------
|
| 125 |
+
# Blocks
|
| 126 |
+
# ---------------------------------------------------------------------------
|
| 127 |
+
|
| 128 |
+
|
| 129 |
+
class PatchEmbed(nn.Module):
|
| 130 |
+
"""Two-stage patch embedding: a low-rank patch projection, then a 1x1 conv."""
|
| 131 |
+
|
| 132 |
+
def __init__(self, patch_size: int, in_channels: int, hidden_size: int, bottleneck: int):
|
| 133 |
+
super().__init__()
|
| 134 |
+
self.proj1 = nn.Conv2d(in_channels, bottleneck, kernel_size=patch_size, stride=patch_size, bias=False)
|
| 135 |
+
self.proj2 = nn.Conv2d(bottleneck, hidden_size, kernel_size=1, bias=True)
|
| 136 |
+
nn.init.xavier_uniform_(self.proj1.weight)
|
| 137 |
+
nn.init.xavier_uniform_(self.proj2.weight)
|
| 138 |
+
nn.init.zeros_(self.proj2.bias)
|
| 139 |
+
|
| 140 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 141 |
+
return self.proj2(self.proj1(x)).flatten(2).transpose(1, 2)
|
| 142 |
+
|
| 143 |
+
|
| 144 |
+
class TextBlock(nn.Module):
|
| 145 |
+
"""Text-only transformer block that refines the T5 tokens before the joint blocks."""
|
| 146 |
+
|
| 147 |
+
def __init__(self, hidden_size: int, num_heads: int, head_dim: int, mlp_ratio: float):
|
| 148 |
+
super().__init__()
|
| 149 |
+
self.num_heads, self.head_dim = num_heads, head_dim
|
| 150 |
+
self.norm1 = RMSNorm(hidden_size)
|
| 151 |
+
self.norm2 = RMSNorm(hidden_size)
|
| 152 |
+
self.qkv = nn.Linear(hidden_size, num_heads * head_dim * 3)
|
| 153 |
+
self.proj = nn.Linear(num_heads * head_dim, hidden_size)
|
| 154 |
+
self.mlp = SwiGLU(hidden_size, int(hidden_size * mlp_ratio))
|
| 155 |
+
self.q_norm = RMSNorm(head_dim)
|
| 156 |
+
self.k_norm = RMSNorm(head_dim)
|
| 157 |
+
|
| 158 |
+
def forward(self, txt: torch.Tensor) -> torch.Tensor:
|
| 159 |
+
b, n, _ = txt.shape
|
| 160 |
+
q, k, v = self.qkv(self.norm1(txt)).view(b, n, 3, self.num_heads, self.head_dim).unbind(2)
|
| 161 |
+
q, k, v = (z.transpose(1, 2) for z in (self.q_norm(q), self.k_norm(k), v))
|
| 162 |
+
out = F.scaled_dot_product_attention(apply_rope(q), apply_rope(k), v, scale=self.head_dim**-0.5)
|
| 163 |
+
txt = txt + self.proj(out.transpose(1, 2).reshape(b, n, -1))
|
| 164 |
+
return txt + self.mlp(self.norm2(txt))
|
| 165 |
+
|
| 166 |
+
|
| 167 |
+
class DoubleStreamBlock(nn.Module):
|
| 168 |
+
"""MMDiT block: separate image/text weights, one joint attention over both.
|
| 169 |
+
|
| 170 |
+
`use_xsa` / `use_attn_gate` turn on self-modulating attention (set only for the
|
| 171 |
+
looped blocks). The attention gate (Qiu et al., 2026:
|
| 172 |
+
https://arxiv.org/abs/2505.06708) is head-wise:
|
| 173 |
+
y_i <- y_i * sigmoid(W_g u_i + b_g), with u_i the block's normed input.
|
| 174 |
+
|
| 175 |
+
`update_text=False` skips the text-stream update, for the last block, whose
|
| 176 |
+
text output is never read.
|
| 177 |
+
"""
|
| 178 |
+
|
| 179 |
+
def __init__(
|
| 180 |
+
self,
|
| 181 |
+
hidden_size: int,
|
| 182 |
+
num_heads: int,
|
| 183 |
+
head_dim: int,
|
| 184 |
+
mlp_ratio: float,
|
| 185 |
+
grid: int,
|
| 186 |
+
use_xsa: bool = False,
|
| 187 |
+
use_attn_gate: bool = False,
|
| 188 |
+
update_text: bool = True,
|
| 189 |
+
):
|
| 190 |
+
super().__init__()
|
| 191 |
+
self.num_heads, self.head_dim, self.grid = num_heads, head_dim, grid
|
| 192 |
+
self.use_xsa, self.use_attn_gate, self.update_text = use_xsa, use_attn_gate, update_text
|
| 193 |
+
inner = num_heads * head_dim
|
| 194 |
+
self.img_norm1 = RMSNorm(hidden_size)
|
| 195 |
+
self.img_norm2 = RMSNorm(hidden_size)
|
| 196 |
+
self.txt_norm1 = RMSNorm(hidden_size)
|
| 197 |
+
self.txt_norm2 = RMSNorm(hidden_size)
|
| 198 |
+
self.img_qkv = nn.Linear(hidden_size, inner * 3)
|
| 199 |
+
self.txt_qkv = nn.Linear(hidden_size, inner * 3)
|
| 200 |
+
self.q_norm = RMSNorm(head_dim)
|
| 201 |
+
self.k_norm = RMSNorm(head_dim)
|
| 202 |
+
self.img_proj = nn.Linear(inner, hidden_size)
|
| 203 |
+
self.txt_proj = nn.Linear(inner, hidden_size)
|
| 204 |
+
if use_attn_gate:
|
| 205 |
+
# Zero bias: the gates start half open on average.
|
| 206 |
+
self.img_gate = nn.Linear(hidden_size, num_heads)
|
| 207 |
+
self.txt_gate = nn.Linear(hidden_size, num_heads)
|
| 208 |
+
nn.init.zeros_(self.img_gate.bias)
|
| 209 |
+
nn.init.zeros_(self.txt_gate.bias)
|
| 210 |
+
self.img_mlp = SwiGLU(hidden_size, int(hidden_size * mlp_ratio))
|
| 211 |
+
self.txt_mlp = SwiGLU(hidden_size, int(hidden_size * mlp_ratio))
|
| 212 |
+
|
| 213 |
+
def forward(self, img: torch.Tensor, txt: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
| 214 |
+
b, li, _ = img.shape
|
| 215 |
+
lt = txt.shape[1]
|
| 216 |
+
img_n, txt_n = self.img_norm1(img), self.txt_norm1(txt)
|
| 217 |
+
qi, ki, vi = self.img_qkv(img_n).view(b, li, 3, self.num_heads, self.head_dim).unbind(2)
|
| 218 |
+
qt, kt, vt = self.txt_qkv(txt_n).view(b, lt, 3, self.num_heads, self.head_dim).unbind(2)
|
| 219 |
+
# Joint sequence [text; image], heads first.
|
| 220 |
+
q = torch.cat([qt, qi], dim=1).transpose(1, 2)
|
| 221 |
+
k = torch.cat([kt, ki], dim=1).transpose(1, 2)
|
| 222 |
+
v = torch.cat([vt, vi], dim=1).transpose(1, 2)
|
| 223 |
+
q = torch.cat([apply_rope(self.q_norm(q[:, :, :lt])), apply_rope(self.q_norm(q[:, :, lt:]), self.grid)], dim=2)
|
| 224 |
+
k = torch.cat([apply_rope(self.k_norm(k[:, :, :lt])), apply_rope(self.k_norm(k[:, :, lt:]), self.grid)], dim=2)
|
| 225 |
+
out = F.scaled_dot_product_attention(q, k, v, scale=self.head_dim**-0.5)
|
| 226 |
+
if self.use_xsa:
|
| 227 |
+
out = exclusive_self_attention(out, v)
|
| 228 |
+
out = out.transpose(1, 2) # [b, tokens, heads, head_dim]
|
| 229 |
+
out_t, out_i = out[:, :lt], out[:, lt:]
|
| 230 |
+
if self.use_attn_gate:
|
| 231 |
+
out_i = out_i * torch.sigmoid(self.img_gate(img_n)).unsqueeze(-1)
|
| 232 |
+
img = img + self.img_proj(out_i.reshape(b, li, -1))
|
| 233 |
+
img = img + self.img_mlp(self.img_norm2(img))
|
| 234 |
+
if self.update_text:
|
| 235 |
+
if self.use_attn_gate:
|
| 236 |
+
out_t = out_t * torch.sigmoid(self.txt_gate(txt_n)).unsqueeze(-1)
|
| 237 |
+
txt = txt + self.txt_proj(out_t.reshape(b, lt, -1))
|
| 238 |
+
txt = txt + self.txt_mlp(self.txt_norm2(txt))
|
| 239 |
+
return img, txt
|
| 240 |
+
|
| 241 |
+
|
| 242 |
+
# ---------------------------------------------------------------------------
|
| 243 |
+
# Looped MMDiT
|
| 244 |
+
# ---------------------------------------------------------------------------
|
| 245 |
+
|
| 246 |
+
|
| 247 |
+
class LoopedDiTTransformer2DModel(ModelMixin, ConfigMixin):
|
| 248 |
+
"""Predicts the clean image x0 from a noisy image and T5 text embeddings.
|
| 249 |
+
|
| 250 |
+
This is a diffusers [`ModelMixin`]: `save_pretrained` / `from_pretrained` round-trip the
|
| 251 |
+
architecture in `config.json`. `loop_split` is `(pre, core, post)` and `num_loops` is the
|
| 252 |
+
trained loop depth (`1` is the MiniT2I model without looping). With `share_loop_weights=False`
|
| 253 |
+
every pass gets its own copy of the core blocks: the compute-matched "deeper" baseline with
|
| 254 |
+
the same exits.
|
| 255 |
+
|
| 256 |
+
The forward pass does not take a timestep. Flow-matching time is applied by the pipeline when
|
| 257 |
+
it converts the x0 prediction into a velocity.
|
| 258 |
+
"""
|
| 259 |
+
|
| 260 |
+
def __init__(
|
| 261 |
+
self,
|
| 262 |
+
image_size: int = 512,
|
| 263 |
+
patch_size: int = 32,
|
| 264 |
+
in_channels: int = 3,
|
| 265 |
+
hidden_size: int = 768,
|
| 266 |
+
num_heads: int = 12,
|
| 267 |
+
head_dim: int = 64,
|
| 268 |
+
mlp_ratio: float = 2.6667,
|
| 269 |
+
pca_channels: int = 128,
|
| 270 |
+
text_dim: int = 1024,
|
| 271 |
+
text_preamble_depth: int = 2,
|
| 272 |
+
loop_split: tuple[int, int, int] | list[int] = (6, 5, 6),
|
| 273 |
+
num_loops: int = 4,
|
| 274 |
+
share_loop_weights: bool = True,
|
| 275 |
+
use_xsa: bool = False,
|
| 276 |
+
use_attn_gate: bool = False,
|
| 277 |
+
):
|
| 278 |
+
super().__init__()
|
| 279 |
+
loop_split = [int(n) for n in loop_split]
|
| 280 |
+
num_loops = int(num_loops)
|
| 281 |
+
share_loop_weights = bool(share_loop_weights)
|
| 282 |
+
use_xsa = bool(use_xsa)
|
| 283 |
+
use_attn_gate = bool(use_attn_gate)
|
| 284 |
+
if len(loop_split) != 3 or min(loop_split) < 1:
|
| 285 |
+
raise ValueError(f"loop_split must be three positive block counts (pre, core, post), got {loop_split}")
|
| 286 |
+
if num_loops < 1:
|
| 287 |
+
raise ValueError(f"num_loops must be >= 1, got {num_loops}")
|
| 288 |
+
# Lists (not tuples) so config.json stays valid JSON.
|
| 289 |
+
self.register_to_config(
|
| 290 |
+
image_size=image_size,
|
| 291 |
+
patch_size=patch_size,
|
| 292 |
+
in_channels=in_channels,
|
| 293 |
+
hidden_size=hidden_size,
|
| 294 |
+
num_heads=num_heads,
|
| 295 |
+
head_dim=head_dim,
|
| 296 |
+
mlp_ratio=mlp_ratio,
|
| 297 |
+
pca_channels=pca_channels,
|
| 298 |
+
text_dim=text_dim,
|
| 299 |
+
text_preamble_depth=text_preamble_depth,
|
| 300 |
+
loop_split=loop_split,
|
| 301 |
+
num_loops=num_loops,
|
| 302 |
+
share_loop_weights=share_loop_weights,
|
| 303 |
+
use_xsa=use_xsa,
|
| 304 |
+
use_attn_gate=use_attn_gate,
|
| 305 |
+
)
|
| 306 |
+
pre, core, post = loop_split
|
| 307 |
+
self.patch_size, self.in_channels = patch_size, in_channels
|
| 308 |
+
self.grid = image_size // patch_size
|
| 309 |
+
self.pre, self.core, self.post = pre, core, post
|
| 310 |
+
self.num_loops = num_loops
|
| 311 |
+
self.share_loop_weights = share_loop_weights
|
| 312 |
+
|
| 313 |
+
self.img_embed = PatchEmbed(patch_size, in_channels, hidden_size, pca_channels)
|
| 314 |
+
self.txt_embed = nn.Linear(text_dim, hidden_size, bias=False)
|
| 315 |
+
# Replaces the T5 embedding at padded prompt positions (and everywhere
|
| 316 |
+
# for the unconditional branch of classifier-free guidance).
|
| 317 |
+
self.mask_token = nn.Parameter(torch.zeros(1, 1, text_dim))
|
| 318 |
+
nn.init.normal_(self.mask_token, std=0.02)
|
| 319 |
+
# MiniT2I's timestep and pooled-text embedders. The model has no timestep
|
| 320 |
+
# conditioning and never uses them; they are kept (frozen, see below) so that
|
| 321 |
+
# the model and its checkpoints match MiniT2I and the paper.
|
| 322 |
+
self.t_embed = nn.ModuleDict(
|
| 323 |
+
{"mlp": nn.Sequential(nn.Linear(256, hidden_size), nn.SiLU(), nn.Linear(hidden_size, hidden_size))}
|
| 324 |
+
)
|
| 325 |
+
for layer in (self.t_embed.mlp[0], self.t_embed.mlp[2]):
|
| 326 |
+
nn.init.normal_(layer.weight, std=0.02)
|
| 327 |
+
nn.init.zeros_(layer.bias)
|
| 328 |
+
self.pooled_embed = nn.Linear(text_dim, hidden_size, bias=False)
|
| 329 |
+
self.register_buffer("pos_embed", sincos_2d(hidden_size, self.grid)[None], persistent=False)
|
| 330 |
+
self.txt_blocks = nn.ModuleList(
|
| 331 |
+
TextBlock(hidden_size, num_heads, head_dim, mlp_ratio) for _ in range(text_preamble_depth)
|
| 332 |
+
)
|
| 333 |
+
looped = core if self.share_loop_weights else core * self.num_loops
|
| 334 |
+
depth = pre + looped + post
|
| 335 |
+
self.blocks = nn.ModuleList(
|
| 336 |
+
DoubleStreamBlock(
|
| 337 |
+
hidden_size,
|
| 338 |
+
num_heads,
|
| 339 |
+
head_dim,
|
| 340 |
+
mlp_ratio,
|
| 341 |
+
self.grid,
|
| 342 |
+
use_xsa=use_xsa and pre <= i < pre + looped,
|
| 343 |
+
use_attn_gate=use_attn_gate and pre <= i < pre + looped,
|
| 344 |
+
update_text=i < depth - 1,
|
| 345 |
+
)
|
| 346 |
+
for i in range(depth)
|
| 347 |
+
)
|
| 348 |
+
self.final_norm = RMSNorm(hidden_size)
|
| 349 |
+
self.final = nn.Linear(hidden_size, patch_size * patch_size * in_channels)
|
| 350 |
+
nn.init.zeros_(self.final.weight)
|
| 351 |
+
nn.init.zeros_(self.final.bias)
|
| 352 |
+
# Frozen because they never receive a gradient: the unused embedders and the
|
| 353 |
+
# text-stream update of the last block (whose text output is never read).
|
| 354 |
+
last = self.blocks[-1]
|
| 355 |
+
for module in (self.t_embed, self.pooled_embed, last.txt_norm2, last.txt_proj, last.txt_mlp):
|
| 356 |
+
module.requires_grad_(False)
|
| 357 |
+
|
| 358 |
+
def unpatchify(self, x: torch.Tensor) -> torch.Tensor:
|
| 359 |
+
b, n, _ = x.shape
|
| 360 |
+
p, c, g = self.patch_size, self.in_channels, int(n**0.5)
|
| 361 |
+
x = x.view(b, g, g, p, p, c).permute(0, 5, 1, 3, 2, 4).contiguous()
|
| 362 |
+
return x.view(b, c, g * p, g * p)
|
| 363 |
+
|
| 364 |
+
def loop_blocks(self, r: int) -> nn.ModuleList:
|
| 365 |
+
"""Core blocks run on loop pass r (1-based)."""
|
| 366 |
+
start = self.pre if self.share_loop_weights else self.pre + (r - 1) * self.core
|
| 367 |
+
return self.blocks[start : start + self.core]
|
| 368 |
+
|
| 369 |
+
def decode(self, img: torch.Tensor, txt: torch.Tensor) -> torch.Tensor:
|
| 370 |
+
"""Post-loop blocks and output head: a loop state -> x0 prediction."""
|
| 371 |
+
for block in self.blocks[len(self.blocks) - self.post :]:
|
| 372 |
+
img, txt = block(img, txt)
|
| 373 |
+
return self.unpatchify(self.final(self.final_norm(img))).float()
|
| 374 |
+
|
| 375 |
+
def forward(
|
| 376 |
+
self,
|
| 377 |
+
x: torch.Tensor,
|
| 378 |
+
text: torch.Tensor,
|
| 379 |
+
text_mask: torch.Tensor,
|
| 380 |
+
num_loops: int | None = None,
|
| 381 |
+
exit_loops: tuple[int, ...] = (),
|
| 382 |
+
) -> torch.Tensor | tuple[torch.Tensor, dict[int, torch.Tensor]]:
|
| 383 |
+
"""x: noisy images [B, C, H, W]; text: T5 states [B, L, text_dim];
|
| 384 |
+
text_mask: [B, L], 1 for prompt tokens (all 0 = unconditional).
|
| 385 |
+
|
| 386 |
+
num_loops overrides the loop depth at inference. exit_loops lists
|
| 387 |
+
intermediate depths r < num_loops to decode as well; the call then
|
| 388 |
+
returns (final prediction, {r: prediction after r loops}).
|
| 389 |
+
"""
|
| 390 |
+
n = self.num_loops if num_loops is None else int(num_loops)
|
| 391 |
+
if n < 1 or (not self.share_loop_weights and n > self.num_loops):
|
| 392 |
+
raise ValueError(f"num_loops={n} is not available for this model (trained with {self.num_loops})")
|
| 393 |
+
exits = sorted({int(r) for r in exit_loops})
|
| 394 |
+
if any(not 1 <= r < n for r in exits):
|
| 395 |
+
raise ValueError(f"exit_loops must lie in [1, {n}), got {exits}")
|
| 396 |
+
|
| 397 |
+
text = torch.where(text_mask.to(torch.bool)[:, :, None], text, self.mask_token.to(text.dtype))
|
| 398 |
+
img = self.img_embed(x) + self.pos_embed.to(device=x.device, dtype=x.dtype)
|
| 399 |
+
txt = self.txt_embed(text)
|
| 400 |
+
for block in self.txt_blocks:
|
| 401 |
+
txt = block(txt)
|
| 402 |
+
for block in self.blocks[: self.pre]:
|
| 403 |
+
img, txt = block(img, txt)
|
| 404 |
+
states = {}
|
| 405 |
+
for r in range(1, n + 1):
|
| 406 |
+
for block in self.loop_blocks(r):
|
| 407 |
+
img, txt = block(img, txt)
|
| 408 |
+
if r in exits:
|
| 409 |
+
states[r] = (img, txt)
|
| 410 |
+
out = self.decode(img, txt)
|
| 411 |
+
if not exits:
|
| 412 |
+
return out
|
| 413 |
+
return out, {r: self.decode(*states[r]) for r in exits}
|
| 414 |
+
|
| 415 |
+
|
| 416 |
+
# Training imports this name. The diffusers class name is the one stored in checkpoints.
|
| 417 |
+
LoopedMMDiT = LoopedDiTTransformer2DModel
|
Looped-DiT-B-32/demo.png
ADDED
|
Git LFS Details
|
Looped-DiT-B-32/model_index.json
ADDED
|
@@ -0,0 +1,23 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"_class_name": [
|
| 3 |
+
"pipeline",
|
| 4 |
+
"LoopedDiTPipeline"
|
| 5 |
+
],
|
| 6 |
+
"_diffusers_version": "0.39.0",
|
| 7 |
+
"scheduler": [
|
| 8 |
+
"diffusers",
|
| 9 |
+
"FlowMatchEulerDiscreteScheduler"
|
| 10 |
+
],
|
| 11 |
+
"text_encoder": [
|
| 12 |
+
"transformers",
|
| 13 |
+
"T5EncoderModel"
|
| 14 |
+
],
|
| 15 |
+
"tokenizer": [
|
| 16 |
+
"transformers",
|
| 17 |
+
"T5Tokenizer"
|
| 18 |
+
],
|
| 19 |
+
"transformer": [
|
| 20 |
+
"transformer_looped_dit",
|
| 21 |
+
"LoopedDiTTransformer2DModel"
|
| 22 |
+
]
|
| 23 |
+
}
|
Looped-DiT-B-32/pipeline.py
ADDED
|
@@ -0,0 +1,763 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2025 The HuggingFace Team. All rights reserved.
|
| 2 |
+
#
|
| 3 |
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
| 4 |
+
# you may not use this file except in compliance with the License.
|
| 5 |
+
# You may obtain a copy of the License at
|
| 6 |
+
#
|
| 7 |
+
# http://www.apache.org/licenses/LICENSE-2.0
|
| 8 |
+
#
|
| 9 |
+
# Unless required by applicable law or agreed to in writing, software
|
| 10 |
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
| 11 |
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
| 12 |
+
# See the License for the specific language governing permissions and
|
| 13 |
+
# limitations under the License.
|
| 14 |
+
|
| 15 |
+
import inspect
|
| 16 |
+
from typing import Any, Callable
|
| 17 |
+
|
| 18 |
+
import torch
|
| 19 |
+
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
|
| 20 |
+
from diffusers.models.modeling_utils import ModelMixin
|
| 21 |
+
from diffusers.pipelines.pipeline_utils import DiffusionPipeline, ImagePipelineOutput
|
| 22 |
+
from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion import retrieve_timesteps
|
| 23 |
+
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler, KarrasDiffusionSchedulers
|
| 24 |
+
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
| 25 |
+
from diffusers.utils import deprecate, is_torch_xla_available, logging, replace_example_docstring
|
| 26 |
+
from diffusers.utils.torch_utils import randn_tensor
|
| 27 |
+
from PIL import Image
|
| 28 |
+
from transformers import AutoTokenizer, T5EncoderModel
|
| 29 |
+
|
| 30 |
+
if is_torch_xla_available():
|
| 31 |
+
import torch_xla.core.xla_model as xm
|
| 32 |
+
|
| 33 |
+
XLA_AVAILABLE = True
|
| 34 |
+
else:
|
| 35 |
+
XLA_AVAILABLE = False
|
| 36 |
+
|
| 37 |
+
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
| 38 |
+
|
| 39 |
+
# Training clamps the flow-matching denominator so the loss stays finite at t -> 1.
|
| 40 |
+
VELOCITY_DENOM_MIN = 0.05
|
| 41 |
+
|
| 42 |
+
EXAMPLE_DOC_STRING = """
|
| 43 |
+
Examples:
|
| 44 |
+
```py
|
| 45 |
+
>>> from pathlib import Path
|
| 46 |
+
>>> import torch
|
| 47 |
+
>>> from diffusers import DiffusionPipeline
|
| 48 |
+
|
| 49 |
+
>>> model_dir = Path("checkpoints/looped-dit-b16").resolve()
|
| 50 |
+
>>> pipe = DiffusionPipeline.from_pretrained(
|
| 51 |
+
... str(model_dir),
|
| 52 |
+
... local_files_only=True,
|
| 53 |
+
... custom_pipeline=str(model_dir / "pipeline.py"),
|
| 54 |
+
... trust_remote_code=True,
|
| 55 |
+
... torch_dtype=torch.bfloat16,
|
| 56 |
+
... ).to("cuda")
|
| 57 |
+
|
| 58 |
+
>>> image = pipe(
|
| 59 |
+
... "a red cube on top of a blue sphere",
|
| 60 |
+
... num_inference_steps=100,
|
| 61 |
+
... guidance_scale=6.0,
|
| 62 |
+
... num_loops=4,
|
| 63 |
+
... generator=torch.Generator(device="cuda").manual_seed(0),
|
| 64 |
+
... ).images[0]
|
| 65 |
+
>>> image.save("sample.png")
|
| 66 |
+
|
| 67 |
+
>>> # Hugging Face Hub style model id: UserID/RepoID
|
| 68 |
+
>>> # RepoID is usually like "modelname-diffusers"
|
| 69 |
+
>>> # Example: "your-user/Looped-DiT-diffusers"
|
| 70 |
+
```
|
| 71 |
+
"""
|
| 72 |
+
|
| 73 |
+
|
| 74 |
+
def paper_euler_sigmas(num_inference_steps: int) -> list[float]:
|
| 75 |
+
r"""
|
| 76 |
+
Sigma grid of the training Euler sampler.
|
| 77 |
+
|
| 78 |
+
Training integrates flow time `t` from 0 (noise) to 1 (data) with
|
| 79 |
+
`torch.linspace(0, 1, steps + 1)`. Flow-match schedulers step in sigma
|
| 80 |
+
`1 - t` and append the terminal 0 themselves, so the returned list omits that 0.
|
| 81 |
+
|
| 82 |
+
Args:
|
| 83 |
+
num_inference_steps (`int`):
|
| 84 |
+
Number of Euler steps. Must be positive.
|
| 85 |
+
|
| 86 |
+
Returns:
|
| 87 |
+
`list[float]`: `num_inference_steps` sigmas starting at 1 and ending at `1 / steps`.
|
| 88 |
+
"""
|
| 89 |
+
if num_inference_steps <= 0:
|
| 90 |
+
raise ValueError(f"`num_inference_steps` must be positive, got {num_inference_steps}.")
|
| 91 |
+
flow_time = torch.linspace(0.0, 1.0, num_inference_steps + 1)
|
| 92 |
+
return (1.0 - flow_time)[:-1].tolist()
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
class LoopedDiTPipeline(DiffusionPipeline):
|
| 96 |
+
r"""
|
| 97 |
+
Text-to-image pipeline for Looped-DiT.
|
| 98 |
+
|
| 99 |
+
Looped-DiT denoises directly in RGB pixel space (no VAE). The transformer predicts the clean
|
| 100 |
+
image `x0`. This pipeline converts that prediction to a flow-matching velocity and integrates it
|
| 101 |
+
with a diffusers scheduler. The default scheduler is [`FlowMatchEulerDiscreteScheduler`] on the
|
| 102 |
+
same uniform grid as the paper (100 steps, shift 1). Any [`KarrasDiffusionSchedulers`] instance
|
| 103 |
+
can be assigned to `pipe.scheduler` without other code changes.
|
| 104 |
+
|
| 105 |
+
Classifier-free guidance uses the training null condition: an all-zero text mask, which the
|
| 106 |
+
denoiser replaces with its mask token. There is no separate negative-prompt encoder.
|
| 107 |
+
|
| 108 |
+
The pipeline inherits from [`DiffusionPipeline`]. Check the superclass documentation for the
|
| 109 |
+
generic methods (download, save, device placement, CPU offload).
|
| 110 |
+
|
| 111 |
+
Args:
|
| 112 |
+
transformer ([`ModelMixin`]):
|
| 113 |
+
Looped-DiT denoiser (`LoopedDiTTransformer2DModel`) that predicts `x0` in pixel space.
|
| 114 |
+
scheduler ([`FlowMatchEulerDiscreteScheduler`] or [`KarrasDiffusionSchedulers`]):
|
| 115 |
+
Scheduler used to step the flow. The paper setting is
|
| 116 |
+
`FlowMatchEulerDiscreteScheduler(num_train_timesteps=1000, shift=1.0)`.
|
| 117 |
+
tokenizer ([`~transformers.AutoTokenizer`], *optional*):
|
| 118 |
+
Tokenizer for the frozen text encoder. Loaded from `text_encoder_name` when missing.
|
| 119 |
+
text_encoder ([`~transformers.T5EncoderModel`], *optional*):
|
| 120 |
+
Frozen FLAN-T5 encoder. Loaded from `text_encoder_name` when missing.
|
| 121 |
+
`DiffusionPipeline.from_pretrained(..., torch_dtype=torch.bfloat16)` keeps this encoder in bf16 with the denoiser.
|
| 122 |
+
text_encoder_name (`str`, defaults to `"google/flan-t5-large"`):
|
| 123 |
+
Hub id or local path used when `tokenizer` / `text_encoder` are not passed.
|
| 124 |
+
prompt_length (`int`, *optional*):
|
| 125 |
+
Token length prompts are padded or truncated to. Defaults to `tokenizer.model_max_length`.
|
| 126 |
+
noise_scale (`float`, defaults to 2.0):
|
| 127 |
+
Standard deviation of the initial Gaussian, matching the training noise scale.
|
| 128 |
+
default_num_inference_steps (`int`, defaults to 100):
|
| 129 |
+
Step count used when `__call__` does not pass `num_inference_steps`.
|
| 130 |
+
"""
|
| 131 |
+
|
| 132 |
+
model_cpu_offload_seq = "text_encoder->transformer"
|
| 133 |
+
_optional_components = ["tokenizer", "text_encoder"]
|
| 134 |
+
_callback_tensor_inputs = ["latents", "prompt_embeds", "prompt_attention_mask"]
|
| 135 |
+
|
| 136 |
+
def __init__(
|
| 137 |
+
self,
|
| 138 |
+
transformer: ModelMixin,
|
| 139 |
+
scheduler: KarrasDiffusionSchedulers | SchedulerMixin,
|
| 140 |
+
tokenizer: Any | None = None,
|
| 141 |
+
text_encoder: T5EncoderModel | None = None,
|
| 142 |
+
text_encoder_name: str = "google/flan-t5-large",
|
| 143 |
+
prompt_length: int | None = None,
|
| 144 |
+
noise_scale: float = 2.0,
|
| 145 |
+
default_num_inference_steps: int = 100,
|
| 146 |
+
):
|
| 147 |
+
super().__init__()
|
| 148 |
+
if prompt_length is None and tokenizer is not None:
|
| 149 |
+
prompt_length = int(getattr(tokenizer, "model_max_length", 256))
|
| 150 |
+
if prompt_length is None:
|
| 151 |
+
prompt_length = 256
|
| 152 |
+
if scheduler is None:
|
| 153 |
+
scheduler = self._default_scheduler()
|
| 154 |
+
if noise_scale <= 0:
|
| 155 |
+
raise ValueError(f"`noise_scale` must be positive, got {noise_scale}.")
|
| 156 |
+
if prompt_length < 1:
|
| 157 |
+
raise ValueError(f"`prompt_length` must be positive, got {prompt_length}.")
|
| 158 |
+
if default_num_inference_steps < 1:
|
| 159 |
+
raise ValueError(f"`default_num_inference_steps` must be positive, got {default_num_inference_steps}.")
|
| 160 |
+
|
| 161 |
+
self.register_modules(
|
| 162 |
+
transformer=transformer,
|
| 163 |
+
scheduler=scheduler,
|
| 164 |
+
tokenizer=tokenizer,
|
| 165 |
+
text_encoder=text_encoder,
|
| 166 |
+
)
|
| 167 |
+
self.register_to_config(
|
| 168 |
+
text_encoder_name=text_encoder_name,
|
| 169 |
+
prompt_length=int(prompt_length),
|
| 170 |
+
noise_scale=float(noise_scale),
|
| 171 |
+
default_num_inference_steps=int(default_num_inference_steps),
|
| 172 |
+
)
|
| 173 |
+
|
| 174 |
+
@staticmethod
|
| 175 |
+
def _default_scheduler() -> FlowMatchEulerDiscreteScheduler:
|
| 176 |
+
r"""
|
| 177 |
+
Build the paper's Euler scheduler.
|
| 178 |
+
|
| 179 |
+
Returns:
|
| 180 |
+
[`FlowMatchEulerDiscreteScheduler`]: 1000 training timesteps, shift 1, deterministic.
|
| 181 |
+
"""
|
| 182 |
+
kwargs: dict[str, Any] = {"num_train_timesteps": 1000, "shift": 1.0}
|
| 183 |
+
if "stochastic_sampling" in inspect.signature(FlowMatchEulerDiscreteScheduler.__init__).parameters:
|
| 184 |
+
kwargs["stochastic_sampling"] = False
|
| 185 |
+
return FlowMatchEulerDiscreteScheduler(**kwargs)
|
| 186 |
+
|
| 187 |
+
def _encode_prompt(
|
| 188 |
+
self,
|
| 189 |
+
prompt: str | list[str] | None,
|
| 190 |
+
device: torch.device,
|
| 191 |
+
num_images_per_prompt: int,
|
| 192 |
+
prompt_embeds: torch.Tensor | None = None,
|
| 193 |
+
prompt_attention_mask: torch.Tensor | None = None,
|
| 194 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 195 |
+
r"""
|
| 196 |
+
Deprecated alias of [`~LoopedDiTPipeline.encode_prompt`].
|
| 197 |
+
|
| 198 |
+
Args:
|
| 199 |
+
prompt (`str` or `list[str]`, *optional*):
|
| 200 |
+
Prompt or prompts to tokenize and encode.
|
| 201 |
+
device (`torch.device`):
|
| 202 |
+
Device of the returned tensors.
|
| 203 |
+
num_images_per_prompt (`int`):
|
| 204 |
+
How many times to repeat each prompt embedding.
|
| 205 |
+
prompt_embeds (`torch.Tensor`, *optional*):
|
| 206 |
+
Already encoded prompts of shape `(batch, sequence, text_dim)`.
|
| 207 |
+
prompt_attention_mask (`torch.Tensor`, *optional*):
|
| 208 |
+
Mask of shape `(batch, sequence)` with 1 on real tokens. Required with `prompt_embeds`
|
| 209 |
+
only when padding should be replaced by the mask token; otherwise a mask of ones is used.
|
| 210 |
+
|
| 211 |
+
Returns:
|
| 212 |
+
`tuple[torch.Tensor, torch.Tensor]`: Prompt embeddings and the attention mask.
|
| 213 |
+
"""
|
| 214 |
+
deprecation_message = (
|
| 215 |
+
"`_encode_prompt()` is deprecated and will be removed in a future version. Use `encode_prompt()` instead."
|
| 216 |
+
)
|
| 217 |
+
deprecate("_encode_prompt()", "1.0.0", deprecation_message, standard_warn=False)
|
| 218 |
+
return self.encode_prompt(prompt, device, num_images_per_prompt, prompt_embeds, prompt_attention_mask)
|
| 219 |
+
|
| 220 |
+
def encode_prompt(
|
| 221 |
+
self,
|
| 222 |
+
prompt: str | list[str] | None,
|
| 223 |
+
device: torch.device,
|
| 224 |
+
num_images_per_prompt: int,
|
| 225 |
+
prompt_embeds: torch.Tensor | None = None,
|
| 226 |
+
prompt_attention_mask: torch.Tensor | None = None,
|
| 227 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 228 |
+
r"""
|
| 229 |
+
Encode prompts with the frozen FLAN-T5 encoder.
|
| 230 |
+
|
| 231 |
+
Prompts are padded or truncated to `config.prompt_length`. The unconditional branch of
|
| 232 |
+
classifier-free guidance is not encoded here: the denoiser builds it by zeroing this mask.
|
| 233 |
+
|
| 234 |
+
Args:
|
| 235 |
+
prompt (`str` or `list[str]`, *optional*):
|
| 236 |
+
Prompt or prompts to tokenize. Ignored when `prompt_embeds` is passed.
|
| 237 |
+
device (`torch.device`):
|
| 238 |
+
Device of the returned tensors.
|
| 239 |
+
num_images_per_prompt (`int`):
|
| 240 |
+
Number of times to repeat each encoded prompt along the batch dimension.
|
| 241 |
+
prompt_embeds (`torch.Tensor`, *optional*):
|
| 242 |
+
Precomputed embeddings of shape `(batch, sequence, text_dim)`. When set, `prompt` is ignored.
|
| 243 |
+
prompt_attention_mask (`torch.Tensor`, *optional*):
|
| 244 |
+
Mask of shape `(batch, sequence)`, 1 for tokens that should condition the model. When
|
| 245 |
+
`prompt_embeds` is set and this is omitted, every position is treated as a real token.
|
| 246 |
+
|
| 247 |
+
Returns:
|
| 248 |
+
`tuple[torch.Tensor, torch.Tensor]`:
|
| 249 |
+
Embeddings `(batch * num_images_per_prompt, sequence, text_dim)` and a mask of the same batch.
|
| 250 |
+
"""
|
| 251 |
+
if num_images_per_prompt < 1:
|
| 252 |
+
raise ValueError(f"`num_images_per_prompt` must be >= 1, got {num_images_per_prompt}.")
|
| 253 |
+
|
| 254 |
+
if prompt_embeds is None:
|
| 255 |
+
if isinstance(prompt, str):
|
| 256 |
+
prompt = [prompt]
|
| 257 |
+
if self.tokenizer is None:
|
| 258 |
+
self.tokenizer = AutoTokenizer.from_pretrained(
|
| 259 |
+
self.config.text_encoder_name, model_max_length=int(self.config.prompt_length)
|
| 260 |
+
)
|
| 261 |
+
if self.text_encoder is None:
|
| 262 |
+
self.text_encoder = T5EncoderModel.from_pretrained(self.config.text_encoder_name)
|
| 263 |
+
self.text_encoder.requires_grad_(False)
|
| 264 |
+
self.text_encoder.eval()
|
| 265 |
+
encoder_device = next(self.text_encoder.parameters()).device
|
| 266 |
+
if encoder_device != device:
|
| 267 |
+
self.text_encoder.to(device)
|
| 268 |
+
tokens = self.tokenizer(
|
| 269 |
+
prompt,
|
| 270 |
+
max_length=int(self.config.prompt_length),
|
| 271 |
+
padding="max_length",
|
| 272 |
+
truncation=True,
|
| 273 |
+
return_tensors="pt",
|
| 274 |
+
)
|
| 275 |
+
input_ids = tokens.input_ids.to(device)
|
| 276 |
+
prompt_attention_mask = tokens.attention_mask.to(device)
|
| 277 |
+
prompt_embeds = self.text_encoder(input_ids=input_ids, attention_mask=prompt_attention_mask).last_hidden_state
|
| 278 |
+
else:
|
| 279 |
+
prompt_embeds = prompt_embeds.to(device)
|
| 280 |
+
if prompt_attention_mask is None:
|
| 281 |
+
prompt_attention_mask = torch.ones(
|
| 282 |
+
prompt_embeds.shape[:2], device=device, dtype=torch.long
|
| 283 |
+
)
|
| 284 |
+
else:
|
| 285 |
+
prompt_attention_mask = prompt_attention_mask.to(device)
|
| 286 |
+
if prompt_embeds.shape[0] != prompt_attention_mask.shape[0]:
|
| 287 |
+
raise ValueError(
|
| 288 |
+
"`prompt_embeds` and `prompt_attention_mask` must have the same batch size, got "
|
| 289 |
+
f"{prompt_embeds.shape[0]} and {prompt_attention_mask.shape[0]}."
|
| 290 |
+
)
|
| 291 |
+
|
| 292 |
+
if num_images_per_prompt != 1:
|
| 293 |
+
prompt_embeds = prompt_embeds.repeat_interleave(num_images_per_prompt, dim=0)
|
| 294 |
+
prompt_attention_mask = prompt_attention_mask.repeat_interleave(num_images_per_prompt, dim=0)
|
| 295 |
+
return prompt_embeds, prompt_attention_mask
|
| 296 |
+
|
| 297 |
+
def prepare_extra_step_kwargs(
|
| 298 |
+
self, generator: torch.Generator | list[torch.Generator] | None, eta: float
|
| 299 |
+
) -> dict[str, Any]:
|
| 300 |
+
r"""
|
| 301 |
+
Extra arguments forwarded to `scheduler.step`, depending on what that method accepts.
|
| 302 |
+
|
| 303 |
+
Args:
|
| 304 |
+
generator (`torch.Generator` or `list[torch.Generator]`, *optional*):
|
| 305 |
+
Generator passed through when the scheduler step samples noise.
|
| 306 |
+
eta (`float`):
|
| 307 |
+
DDIM eta in `[0, 1]`. Ignored by schedulers whose `step` has no `eta` argument.
|
| 308 |
+
|
| 309 |
+
Returns:
|
| 310 |
+
`dict`: Keyword arguments for `scheduler.step`.
|
| 311 |
+
"""
|
| 312 |
+
extra_step_kwargs: dict[str, Any] = {}
|
| 313 |
+
step_params = set(inspect.signature(self.scheduler.step).parameters.keys())
|
| 314 |
+
if "eta" in step_params:
|
| 315 |
+
extra_step_kwargs["eta"] = eta
|
| 316 |
+
if "generator" in step_params:
|
| 317 |
+
extra_step_kwargs["generator"] = generator
|
| 318 |
+
return extra_step_kwargs
|
| 319 |
+
|
| 320 |
+
def check_inputs(
|
| 321 |
+
self,
|
| 322 |
+
prompt: str | list[str] | None,
|
| 323 |
+
height: int,
|
| 324 |
+
width: int,
|
| 325 |
+
callback_steps: int | None,
|
| 326 |
+
prompt_embeds: torch.Tensor | None = None,
|
| 327 |
+
prompt_attention_mask: torch.Tensor | None = None,
|
| 328 |
+
callback_on_step_end_tensor_inputs: list[str] | None = None,
|
| 329 |
+
num_inference_steps: int = 100,
|
| 330 |
+
guidance_scale: float = 6.0,
|
| 331 |
+
num_loops: int | None = None,
|
| 332 |
+
output_type: str = "pil",
|
| 333 |
+
) -> None:
|
| 334 |
+
r"""
|
| 335 |
+
Validate generation arguments and raise `ValueError` or `TypeError` on misuse.
|
| 336 |
+
|
| 337 |
+
Args:
|
| 338 |
+
prompt (`str` or `list[str]`, *optional*):
|
| 339 |
+
Prompt text. Mutually exclusive with `prompt_embeds`.
|
| 340 |
+
height (`int`):
|
| 341 |
+
Output height in pixels. Must equal the transformer's trained `image_size`.
|
| 342 |
+
width (`int`):
|
| 343 |
+
Output width in pixels. Must equal the transformer's trained `image_size`.
|
| 344 |
+
callback_steps (`int`, *optional*):
|
| 345 |
+
Deprecated callback period. When set, it must be a positive integer.
|
| 346 |
+
prompt_embeds (`torch.Tensor`, *optional*):
|
| 347 |
+
Precomputed text embeddings. Required when `prompt` is omitted.
|
| 348 |
+
prompt_attention_mask (`torch.Tensor`, *optional*):
|
| 349 |
+
Mask paired with `prompt_embeds`.
|
| 350 |
+
callback_on_step_end_tensor_inputs (`list[str]`, *optional*):
|
| 351 |
+
Tensor names the step callback may read. Each name must be listed on
|
| 352 |
+
`_callback_tensor_inputs`.
|
| 353 |
+
num_inference_steps (`int`):
|
| 354 |
+
Denoising steps. Must be positive.
|
| 355 |
+
guidance_scale (`float`):
|
| 356 |
+
Classifier-free guidance scale. Must be finite. `1` disables guidance.
|
| 357 |
+
num_loops (`int`, *optional*):
|
| 358 |
+
Loop depth. `None` uses the depth stored on the transformer. Otherwise `>= 1`, and
|
| 359 |
+
untied models cannot exceed the trained depth.
|
| 360 |
+
output_type (`str`):
|
| 361 |
+
One of `"pil"`, `"np"`, `"pt"`, or `"latent"`.
|
| 362 |
+
"""
|
| 363 |
+
image_size = int(self.transformer.config.image_size)
|
| 364 |
+
patch_size = int(self.transformer.config.patch_size)
|
| 365 |
+
if height != image_size or width != image_size:
|
| 366 |
+
raise ValueError(
|
| 367 |
+
f"Looped-DiT uses a fixed positional grid of {image_size}x{image_size} "
|
| 368 |
+
f"(patch size {patch_size}). Got height={height}, width={width}."
|
| 369 |
+
)
|
| 370 |
+
if height % patch_size != 0 or width % patch_size != 0:
|
| 371 |
+
raise ValueError(f"height and width must be divisible by patch_size={patch_size}, got {(height, width)}.")
|
| 372 |
+
|
| 373 |
+
if callback_steps is not None and (not isinstance(callback_steps, int) or callback_steps <= 0):
|
| 374 |
+
raise ValueError(
|
| 375 |
+
f"`callback_steps` has to be a positive integer but is {callback_steps} of type {type(callback_steps)}."
|
| 376 |
+
)
|
| 377 |
+
if callback_on_step_end_tensor_inputs is not None and not all(
|
| 378 |
+
key in self._callback_tensor_inputs for key in callback_on_step_end_tensor_inputs
|
| 379 |
+
):
|
| 380 |
+
unexpected = [key for key in callback_on_step_end_tensor_inputs if key not in self._callback_tensor_inputs]
|
| 381 |
+
raise ValueError(
|
| 382 |
+
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {unexpected}."
|
| 383 |
+
)
|
| 384 |
+
|
| 385 |
+
if prompt is not None and prompt_embeds is not None:
|
| 386 |
+
raise ValueError("Cannot forward both `prompt` and `prompt_embeds`. Pass only one of them.")
|
| 387 |
+
if prompt is None and prompt_embeds is None:
|
| 388 |
+
raise ValueError("Provide either `prompt` or `prompt_embeds`.")
|
| 389 |
+
if prompt is not None and not isinstance(prompt, str) and not (
|
| 390 |
+
isinstance(prompt, list) and all(isinstance(item, str) for item in prompt)
|
| 391 |
+
):
|
| 392 |
+
raise TypeError(f"`prompt` has to be a string or a list of strings, got {type(prompt)}.")
|
| 393 |
+
if prompt_embeds is not None and prompt_embeds.ndim != 3:
|
| 394 |
+
raise ValueError(f"`prompt_embeds` must have shape (batch, sequence, dim), got {tuple(prompt_embeds.shape)}.")
|
| 395 |
+
if prompt_attention_mask is not None and prompt_embeds is None:
|
| 396 |
+
raise ValueError("`prompt_attention_mask` was passed without `prompt_embeds`.")
|
| 397 |
+
|
| 398 |
+
if num_inference_steps <= 0:
|
| 399 |
+
raise ValueError(f"`num_inference_steps` must be positive, got {num_inference_steps}.")
|
| 400 |
+
if not torch.isfinite(torch.tensor(guidance_scale)):
|
| 401 |
+
raise ValueError(f"`guidance_scale` must be finite, got {guidance_scale}.")
|
| 402 |
+
if num_loops is not None:
|
| 403 |
+
if int(num_loops) < 1:
|
| 404 |
+
raise ValueError(f"`num_loops` must be >= 1, got {num_loops}.")
|
| 405 |
+
trained = int(self.transformer.config.num_loops)
|
| 406 |
+
if not bool(self.transformer.config.share_loop_weights) and int(num_loops) > trained:
|
| 407 |
+
raise ValueError(
|
| 408 |
+
f"This checkpoint does not share loop weights, so `num_loops` cannot exceed the trained "
|
| 409 |
+
f"depth {trained}. Got {num_loops}."
|
| 410 |
+
)
|
| 411 |
+
if output_type not in {"pil", "np", "pt", "latent"}:
|
| 412 |
+
raise ValueError(f"Unsupported `output_type` {output_type!r}. Choose from 'pil', 'np', 'pt', 'latent'.")
|
| 413 |
+
|
| 414 |
+
def prepare_latents(
|
| 415 |
+
self,
|
| 416 |
+
batch_size: int,
|
| 417 |
+
num_channels: int,
|
| 418 |
+
height: int,
|
| 419 |
+
width: int,
|
| 420 |
+
dtype: torch.dtype,
|
| 421 |
+
device: torch.device,
|
| 422 |
+
generator: torch.Generator | list[torch.Generator] | None,
|
| 423 |
+
latents: torch.Tensor | None = None,
|
| 424 |
+
) -> torch.Tensor:
|
| 425 |
+
r"""
|
| 426 |
+
Sample the initial pixel-space noise, or validate a tensor the caller already sampled.
|
| 427 |
+
|
| 428 |
+
Looped-DiT has no VAE, so "latents" here are RGB images. Fresh noise is scaled by
|
| 429 |
+
`config.noise_scale` (2.0 in the paper). A provided `latents` tensor is not rescaled.
|
| 430 |
+
|
| 431 |
+
Args:
|
| 432 |
+
batch_size (`int`):
|
| 433 |
+
Number of images, including `num_images_per_prompt`.
|
| 434 |
+
num_channels (`int`):
|
| 435 |
+
Channel count. 3 for RGB.
|
| 436 |
+
height (`int`):
|
| 437 |
+
Image height in pixels.
|
| 438 |
+
width (`int`):
|
| 439 |
+
Image width in pixels.
|
| 440 |
+
dtype (`torch.dtype`):
|
| 441 |
+
Dtype of freshly sampled noise. The integration itself is accumulated in float32.
|
| 442 |
+
device (`torch.device`):
|
| 443 |
+
Device of the returned tensor.
|
| 444 |
+
generator (`torch.Generator` or `list[torch.Generator]`, *optional*):
|
| 445 |
+
Per-call RNG. A list must have length `batch_size`.
|
| 446 |
+
latents (`torch.Tensor`, *optional*):
|
| 447 |
+
Starting noise of shape `(batch_size, num_channels, height, width)`.
|
| 448 |
+
|
| 449 |
+
Returns:
|
| 450 |
+
`torch.Tensor`: Starting noise of shape `(batch_size, num_channels, height, width)`.
|
| 451 |
+
"""
|
| 452 |
+
shape = (batch_size, num_channels, height, width)
|
| 453 |
+
if isinstance(generator, list) and len(generator) != batch_size:
|
| 454 |
+
raise ValueError(
|
| 455 |
+
f"You passed a list of {len(generator)} generators for a batch of {batch_size}. "
|
| 456 |
+
"The two lengths must match."
|
| 457 |
+
)
|
| 458 |
+
if latents is None:
|
| 459 |
+
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
| 460 |
+
latents = latents * float(self.config.noise_scale)
|
| 461 |
+
else:
|
| 462 |
+
latents = latents.to(device=device)
|
| 463 |
+
if tuple(latents.shape) != shape:
|
| 464 |
+
raise ValueError(f"`latents` shape {tuple(latents.shape)} does not match the expected {shape}.")
|
| 465 |
+
return latents
|
| 466 |
+
|
| 467 |
+
@property
|
| 468 |
+
def guidance_scale(self) -> float:
|
| 469 |
+
r"""
|
| 470 |
+
Classifier-free guidance scale of the call that is currently running.
|
| 471 |
+
|
| 472 |
+
Returns:
|
| 473 |
+
`float`: The scale set by the active `__call__`.
|
| 474 |
+
"""
|
| 475 |
+
return self._guidance_scale
|
| 476 |
+
|
| 477 |
+
@property
|
| 478 |
+
def do_classifier_free_guidance(self) -> bool:
|
| 479 |
+
r"""
|
| 480 |
+
Whether the active call runs a conditional and an unconditional forward.
|
| 481 |
+
|
| 482 |
+
Returns:
|
| 483 |
+
`bool`: True when `guidance_scale != 1`.
|
| 484 |
+
"""
|
| 485 |
+
return self._guidance_scale != 1.0
|
| 486 |
+
|
| 487 |
+
@property
|
| 488 |
+
def num_timesteps(self) -> int:
|
| 489 |
+
r"""
|
| 490 |
+
Number of scheduler timesteps in the active call.
|
| 491 |
+
|
| 492 |
+
Returns:
|
| 493 |
+
`int`: Length of the timestep schedule.
|
| 494 |
+
"""
|
| 495 |
+
return self._num_timesteps
|
| 496 |
+
|
| 497 |
+
@property
|
| 498 |
+
def interrupt(self) -> bool:
|
| 499 |
+
r"""
|
| 500 |
+
Whether the active denoising loop should skip remaining steps.
|
| 501 |
+
|
| 502 |
+
Returns:
|
| 503 |
+
`bool`: True after the caller sets `pipeline._interrupt = True`.
|
| 504 |
+
"""
|
| 505 |
+
return self._interrupt
|
| 506 |
+
|
| 507 |
+
def _images_from_latents(self, latents: torch.Tensor, output_type: str) -> torch.Tensor | list[Image.Image] | Any:
|
| 508 |
+
r"""
|
| 509 |
+
Convert pixel-space samples in `[-1, 1]` to the requested output type.
|
| 510 |
+
|
| 511 |
+
Quantization matches the original sampler: `uint8(clamp(x, -1, 1) * 127.5 + 128)`.
|
| 512 |
+
|
| 513 |
+
Args:
|
| 514 |
+
latents (`torch.Tensor`):
|
| 515 |
+
Samples of shape `(batch, channels, height, width)` in model range `[-1, 1]`.
|
| 516 |
+
output_type (`str`):
|
| 517 |
+
`"latent"` returns `latents` unchanged. `"pt"` is float RGB in `[0, 1]`. `"np"` is
|
| 518 |
+
`uint8` HWC arrays. `"pil"` is a list of `PIL.Image.Image`.
|
| 519 |
+
|
| 520 |
+
Returns:
|
| 521 |
+
Images in the requested type.
|
| 522 |
+
"""
|
| 523 |
+
if output_type == "latent":
|
| 524 |
+
return latents
|
| 525 |
+
images = (latents.float().clamp(-1, 1) * 127.5 + 128.0).clamp(0, 255).to(torch.uint8)
|
| 526 |
+
if output_type == "pt":
|
| 527 |
+
return images.float() / 255.0
|
| 528 |
+
arrays = images.permute(0, 2, 3, 1).cpu().numpy()
|
| 529 |
+
if output_type == "np":
|
| 530 |
+
return arrays
|
| 531 |
+
return [Image.fromarray(image) for image in arrays]
|
| 532 |
+
|
| 533 |
+
@torch.no_grad()
|
| 534 |
+
@replace_example_docstring(EXAMPLE_DOC_STRING)
|
| 535 |
+
def __call__(
|
| 536 |
+
self,
|
| 537 |
+
prompt: str | list[str] | None = None,
|
| 538 |
+
height: int | None = None,
|
| 539 |
+
width: int | None = None,
|
| 540 |
+
num_inference_steps: int | None = None,
|
| 541 |
+
timesteps: list[int] | None = None,
|
| 542 |
+
sigmas: list[float] | None = None,
|
| 543 |
+
guidance_scale: float = 6.0,
|
| 544 |
+
num_images_per_prompt: int = 1,
|
| 545 |
+
num_loops: int | None = None,
|
| 546 |
+
eta: float = 0.0,
|
| 547 |
+
generator: torch.Generator | list[torch.Generator] | None = None,
|
| 548 |
+
latents: torch.Tensor | None = None,
|
| 549 |
+
prompt_embeds: torch.Tensor | None = None,
|
| 550 |
+
prompt_attention_mask: torch.Tensor | None = None,
|
| 551 |
+
output_type: str = "pil",
|
| 552 |
+
return_dict: bool = True,
|
| 553 |
+
callback_on_step_end: Callable[[int, int, dict], dict] | PipelineCallback | MultiPipelineCallbacks | None = None,
|
| 554 |
+
callback_on_step_end_tensor_inputs: list[str] = ["latents"],
|
| 555 |
+
**kwargs,
|
| 556 |
+
) -> ImagePipelineOutput | tuple:
|
| 557 |
+
r"""
|
| 558 |
+
Generate images from text prompts.
|
| 559 |
+
|
| 560 |
+
Args:
|
| 561 |
+
prompt (`str` or `list[str]`, *optional*):
|
| 562 |
+
Prompt or prompts to guide image generation. Required unless `prompt_embeds` is passed.
|
| 563 |
+
height (`int`, *optional*):
|
| 564 |
+
Image height in pixels. Defaults to the transformer's trained resolution (512).
|
| 565 |
+
Other resolutions are rejected: the positional embedding is a fixed grid.
|
| 566 |
+
width (`int`, *optional*):
|
| 567 |
+
Image width in pixels. Defaults to the trained resolution and must match `height`.
|
| 568 |
+
num_inference_steps (`int`, *optional*):
|
| 569 |
+
Denoising steps. Defaults to `config.default_num_inference_steps` (100).
|
| 570 |
+
timesteps (`list[int]`, *optional*):
|
| 571 |
+
Custom scheduler timesteps, descending. Mutually exclusive with `sigmas`. Ignored by the
|
| 572 |
+
paper Euler grid, which is selected only when both `timesteps` and `sigmas` are omitted
|
| 573 |
+
and the scheduler is [`FlowMatchEulerDiscreteScheduler`].
|
| 574 |
+
sigmas (`list[float]`, *optional*):
|
| 575 |
+
Custom sigmas passed to `scheduler.set_timesteps`. Mutually exclusive with `timesteps`.
|
| 576 |
+
guidance_scale (`float`, defaults to 6.0):
|
| 577 |
+
Classifier-free guidance scale from the paper. Guidance is on when this is not `1`.
|
| 578 |
+
The unconditional branch is an empty text mask, not a negative prompt.
|
| 579 |
+
num_images_per_prompt (`int`, defaults to 1):
|
| 580 |
+
How many images to sample for each prompt.
|
| 581 |
+
num_loops (`int`, *optional*):
|
| 582 |
+
How many times to run the shared middle blocks. `None` uses the trained depth. Other
|
| 583 |
+
depths work without retraining when loop weights are shared.
|
| 584 |
+
eta (`float`, defaults to 0.0):
|
| 585 |
+
DDIM eta. Ignored by the flow-match Euler scheduler.
|
| 586 |
+
generator (`torch.Generator` or `list[torch.Generator]`, *optional*):
|
| 587 |
+
RNG for the initial noise. `None` uses PyTorch's global generator, which is what
|
| 588 |
+
`torch.manual_seed` seeds.
|
| 589 |
+
latents (`torch.Tensor`, *optional*):
|
| 590 |
+
Initial noise `(batch, 3, height, width)`. Not multiplied by `noise_scale`.
|
| 591 |
+
prompt_embeds (`torch.Tensor`, *optional*):
|
| 592 |
+
Precomputed FLAN-T5 states `(batch, sequence, text_dim)` in place of `prompt`.
|
| 593 |
+
prompt_attention_mask (`torch.Tensor`, *optional*):
|
| 594 |
+
Mask `(batch, sequence)` paired with `prompt_embeds`. 1 marks real tokens.
|
| 595 |
+
output_type (`str`, defaults to `"pil"`):
|
| 596 |
+
`"pil"`, `"np"`, `"pt"` (float RGB in `[0, 1]`), or `"latent"` (pixels in model range).
|
| 597 |
+
return_dict (`bool`, defaults to `True`):
|
| 598 |
+
Return [`ImagePipelineOutput`] when `True`, otherwise a one-tuple of images.
|
| 599 |
+
callback_on_step_end (`Callable` or `PipelineCallback`, *optional*):
|
| 600 |
+
Called as `callback_on_step_end(pipeline, step, timestep, callback_kwargs)` after each
|
| 601 |
+
scheduler step. Return a dict to replace tensors listed in
|
| 602 |
+
`callback_on_step_end_tensor_inputs`.
|
| 603 |
+
callback_on_step_end_tensor_inputs (`list[str]`, defaults to `["latents"]`):
|
| 604 |
+
Tensor names passed to the step callback. Must be a subset of `_callback_tensor_inputs`.
|
| 605 |
+
|
| 606 |
+
Examples:
|
| 607 |
+
|
| 608 |
+
Returns:
|
| 609 |
+
[`ImagePipelineOutput`] or `tuple`:
|
| 610 |
+
When `return_dict` is `True`, [`ImagePipelineOutput`] with the images. Otherwise a tuple
|
| 611 |
+
whose first element is the images.
|
| 612 |
+
"""
|
| 613 |
+
callback = kwargs.pop("callback", None)
|
| 614 |
+
callback_steps = kwargs.pop("callback_steps", None)
|
| 615 |
+
if kwargs:
|
| 616 |
+
raise TypeError(f"Unexpected arguments: {sorted(kwargs)}.")
|
| 617 |
+
|
| 618 |
+
if callback is not None:
|
| 619 |
+
deprecate(
|
| 620 |
+
"callback",
|
| 621 |
+
"1.0.0",
|
| 622 |
+
"Passing `callback` as an input argument to `__call__` is deprecated, consider using `callback_on_step_end`",
|
| 623 |
+
)
|
| 624 |
+
if callback_steps is not None:
|
| 625 |
+
deprecate(
|
| 626 |
+
"callback_steps",
|
| 627 |
+
"1.0.0",
|
| 628 |
+
"Passing `callback_steps` as an input argument to `__call__` is deprecated, consider using `callback_on_step_end`",
|
| 629 |
+
)
|
| 630 |
+
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
|
| 631 |
+
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
|
| 632 |
+
|
| 633 |
+
image_size = int(self.transformer.config.image_size)
|
| 634 |
+
height = image_size if height is None else int(height)
|
| 635 |
+
width = image_size if width is None else int(width)
|
| 636 |
+
if num_inference_steps is None:
|
| 637 |
+
num_inference_steps = int(self.config.default_num_inference_steps)
|
| 638 |
+
|
| 639 |
+
# 1. Check inputs.
|
| 640 |
+
self.check_inputs(
|
| 641 |
+
prompt,
|
| 642 |
+
height,
|
| 643 |
+
width,
|
| 644 |
+
callback_steps,
|
| 645 |
+
prompt_embeds,
|
| 646 |
+
prompt_attention_mask,
|
| 647 |
+
callback_on_step_end_tensor_inputs,
|
| 648 |
+
num_inference_steps,
|
| 649 |
+
guidance_scale,
|
| 650 |
+
num_loops,
|
| 651 |
+
output_type,
|
| 652 |
+
)
|
| 653 |
+
|
| 654 |
+
self._guidance_scale = float(guidance_scale)
|
| 655 |
+
self._interrupt = False
|
| 656 |
+
|
| 657 |
+
# 2. Define call parameters.
|
| 658 |
+
if prompt is not None and isinstance(prompt, str):
|
| 659 |
+
batch_size = 1
|
| 660 |
+
elif prompt is not None and isinstance(prompt, list):
|
| 661 |
+
batch_size = len(prompt)
|
| 662 |
+
else:
|
| 663 |
+
batch_size = prompt_embeds.shape[0]
|
| 664 |
+
device = self._execution_device
|
| 665 |
+
|
| 666 |
+
# 3. Encode input prompt.
|
| 667 |
+
prompt_embeds, prompt_attention_mask = self.encode_prompt(
|
| 668 |
+
prompt,
|
| 669 |
+
device,
|
| 670 |
+
num_images_per_prompt,
|
| 671 |
+
prompt_embeds=prompt_embeds,
|
| 672 |
+
prompt_attention_mask=prompt_attention_mask,
|
| 673 |
+
)
|
| 674 |
+
prompt_embeds = prompt_embeds.to(device=device, dtype=self.transformer.dtype)
|
| 675 |
+
prompt_attention_mask = prompt_attention_mask.to(device=device)
|
| 676 |
+
|
| 677 |
+
# 4. Prepare timesteps.
|
| 678 |
+
# The paper's Euler grid is the training linspace. Other schedulers keep their own spacing.
|
| 679 |
+
# A caller-supplied `timesteps` or `sigmas` always wins.
|
| 680 |
+
if timesteps is None and sigmas is None and isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler):
|
| 681 |
+
sigmas = paper_euler_sigmas(num_inference_steps)
|
| 682 |
+
timesteps, num_inference_steps = retrieve_timesteps(self.scheduler, num_inference_steps, device, timesteps, sigmas)
|
| 683 |
+
if getattr(self.scheduler.config, "stochastic_sampling", False):
|
| 684 |
+
raise ValueError(
|
| 685 |
+
"Looped-DiT's training sampler is deterministic. Set `stochastic_sampling=False` on "
|
| 686 |
+
"FlowMatchEulerDiscreteScheduler, or assign a different scheduler."
|
| 687 |
+
)
|
| 688 |
+
|
| 689 |
+
# 5. Prepare latent variables (pixel-space noise; there is no VAE).
|
| 690 |
+
latents = self.prepare_latents(
|
| 691 |
+
batch_size * num_images_per_prompt,
|
| 692 |
+
int(self.transformer.config.in_channels),
|
| 693 |
+
height,
|
| 694 |
+
width,
|
| 695 |
+
self.transformer.dtype,
|
| 696 |
+
device,
|
| 697 |
+
generator,
|
| 698 |
+
latents,
|
| 699 |
+
)
|
| 700 |
+
|
| 701 |
+
# 6. Prepare extra step kwargs.
|
| 702 |
+
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
|
| 703 |
+
num_train_timesteps = int(self.scheduler.config.num_train_timesteps)
|
| 704 |
+
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
| 705 |
+
self._num_timesteps = len(timesteps)
|
| 706 |
+
|
| 707 |
+
# 7. Denoising loop.
|
| 708 |
+
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
| 709 |
+
for i, t in enumerate(timesteps):
|
| 710 |
+
if self.interrupt:
|
| 711 |
+
continue
|
| 712 |
+
|
| 713 |
+
model_latents = latents
|
| 714 |
+
if hasattr(self.scheduler, "scale_model_input"):
|
| 715 |
+
model_latents = self.scheduler.scale_model_input(model_latents, t)
|
| 716 |
+
text = prompt_embeds
|
| 717 |
+
mask = prompt_attention_mask
|
| 718 |
+
if self.do_classifier_free_guidance:
|
| 719 |
+
model_latents = torch.cat([model_latents, model_latents], dim=0)
|
| 720 |
+
text = torch.cat([text, text], dim=0)
|
| 721 |
+
mask = torch.cat([mask, torch.zeros_like(mask)], dim=0)
|
| 722 |
+
|
| 723 |
+
# fp32 latents with bf16 weights match the old sampler, which autocasts the forward.
|
| 724 |
+
amp_dtype = self.transformer.dtype
|
| 725 |
+
use_amp = model_latents.is_cuda and amp_dtype in (torch.float16, torch.bfloat16)
|
| 726 |
+
if not use_amp:
|
| 727 |
+
model_latents = model_latents.to(dtype=amp_dtype)
|
| 728 |
+
with torch.autocast("cuda", dtype=amp_dtype, enabled=use_amp):
|
| 729 |
+
x0 = self.transformer(model_latents, text, mask, num_loops=num_loops)
|
| 730 |
+
x0 = x0.float()
|
| 731 |
+
if self.do_classifier_free_guidance:
|
| 732 |
+
x0_cond, x0_uncond = x0.chunk(2)
|
| 733 |
+
x0 = x0_uncond + self.guidance_scale * (x0_cond - x0_uncond)
|
| 734 |
+
|
| 735 |
+
# sigma = 1 - t_flow. Passing -velocity makes `x + (sigma_next - sigma) * model_output`
|
| 736 |
+
# equal the training update `x + velocity * (t_next - t)`.
|
| 737 |
+
flow_time = (1.0 - t.to(device=latents.device, dtype=torch.float32) / num_train_timesteps)
|
| 738 |
+
velocity = (x0 - latents.float()) / (1.0 - flow_time).clamp_min(VELOCITY_DENOM_MIN)
|
| 739 |
+
latents = self.scheduler.step(-velocity, t, latents, **extra_step_kwargs, return_dict=False)[0]
|
| 740 |
+
|
| 741 |
+
if callback_on_step_end is not None:
|
| 742 |
+
callback_kwargs = {key: locals()[key] for key in callback_on_step_end_tensor_inputs}
|
| 743 |
+
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
|
| 744 |
+
latents = callback_outputs.pop("latents", latents)
|
| 745 |
+
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
|
| 746 |
+
prompt_attention_mask = callback_outputs.pop("prompt_attention_mask", prompt_attention_mask)
|
| 747 |
+
|
| 748 |
+
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
| 749 |
+
progress_bar.update()
|
| 750 |
+
if callback is not None and i % callback_steps == 0:
|
| 751 |
+
callback(i, t, latents)
|
| 752 |
+
|
| 753 |
+
if XLA_AVAILABLE:
|
| 754 |
+
xm.mark_step()
|
| 755 |
+
|
| 756 |
+
images = self._images_from_latents(latents, output_type)
|
| 757 |
+
|
| 758 |
+
# Offload all models.
|
| 759 |
+
self.maybe_free_model_hooks()
|
| 760 |
+
|
| 761 |
+
if not return_dict:
|
| 762 |
+
return (images,)
|
| 763 |
+
return ImagePipelineOutput(images=images)
|
Looped-DiT-B-32/scheduler/scheduler_config.json
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"_class_name": "FlowMatchEulerDiscreteScheduler",
|
| 3 |
+
"_diffusers_version": "0.39.0",
|
| 4 |
+
"base_image_seq_len": 256,
|
| 5 |
+
"base_shift": 0.5,
|
| 6 |
+
"invert_sigmas": false,
|
| 7 |
+
"max_image_seq_len": 4096,
|
| 8 |
+
"max_shift": 1.15,
|
| 9 |
+
"num_train_timesteps": 1000,
|
| 10 |
+
"shift": 1.0,
|
| 11 |
+
"shift_terminal": null,
|
| 12 |
+
"stochastic_sampling": false,
|
| 13 |
+
"time_shift_type": "exponential",
|
| 14 |
+
"use_beta_sigmas": false,
|
| 15 |
+
"use_dynamic_shifting": false,
|
| 16 |
+
"use_exponential_sigmas": false,
|
| 17 |
+
"use_karras_sigmas": false
|
| 18 |
+
}
|
Looped-DiT-B-32/text_encoder/README.md
ADDED
|
@@ -0,0 +1,276 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
language:
|
| 3 |
+
- en
|
| 4 |
+
- fr
|
| 5 |
+
- ro
|
| 6 |
+
- de
|
| 7 |
+
- multilingual
|
| 8 |
+
|
| 9 |
+
widget:
|
| 10 |
+
- text: "Translate to German: My name is Arthur"
|
| 11 |
+
example_title: "Translation"
|
| 12 |
+
- text: "Please answer to the following question. Who is going to be the next Ballon d'or?"
|
| 13 |
+
example_title: "Question Answering"
|
| 14 |
+
- text: "Q: Can Geoffrey Hinton have a conversation with George Washington? Give the rationale before answering."
|
| 15 |
+
example_title: "Logical reasoning"
|
| 16 |
+
- text: "Please answer the following question. What is the boiling point of Nitrogen?"
|
| 17 |
+
example_title: "Scientific knowledge"
|
| 18 |
+
- text: "Answer the following yes/no question. Can you write a whole Haiku in a single tweet?"
|
| 19 |
+
example_title: "Yes/no question"
|
| 20 |
+
- text: "Answer the following yes/no question by reasoning step-by-step. Can you write a whole Haiku in a single tweet?"
|
| 21 |
+
example_title: "Reasoning task"
|
| 22 |
+
- text: "Q: ( False or not False or False ) is? A: Let's think step by step"
|
| 23 |
+
example_title: "Boolean Expressions"
|
| 24 |
+
- text: "The square root of x is the cube root of y. What is y to the power of 2, if x = 4?"
|
| 25 |
+
example_title: "Math reasoning"
|
| 26 |
+
- text: "Premise: At my age you will probably have learnt one lesson. Hypothesis: It's not certain how many lessons you'll learn by your thirties. Does the premise entail the hypothesis?"
|
| 27 |
+
example_title: "Premise and hypothesis"
|
| 28 |
+
|
| 29 |
+
tags:
|
| 30 |
+
- text2text-generation
|
| 31 |
+
|
| 32 |
+
datasets:
|
| 33 |
+
- svakulenk0/qrecc
|
| 34 |
+
- taskmaster2
|
| 35 |
+
- djaym7/wiki_dialog
|
| 36 |
+
- deepmind/code_contests
|
| 37 |
+
- lambada
|
| 38 |
+
- gsm8k
|
| 39 |
+
- aqua_rat
|
| 40 |
+
- esnli
|
| 41 |
+
- quasc
|
| 42 |
+
- qed
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
license: apache-2.0
|
| 46 |
+
---
|
| 47 |
+
|
| 48 |
+
# Model Card for FLAN-T5 large
|
| 49 |
+
|
| 50 |
+
<img src="https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/transformers/model_doc/flan2_architecture.jpg"
|
| 51 |
+
alt="drawing" width="600"/>
|
| 52 |
+
|
| 53 |
+
# Table of Contents
|
| 54 |
+
|
| 55 |
+
0. [TL;DR](#TL;DR)
|
| 56 |
+
1. [Model Details](#model-details)
|
| 57 |
+
2. [Usage](#usage)
|
| 58 |
+
3. [Uses](#uses)
|
| 59 |
+
4. [Bias, Risks, and Limitations](#bias-risks-and-limitations)
|
| 60 |
+
5. [Training Details](#training-details)
|
| 61 |
+
6. [Evaluation](#evaluation)
|
| 62 |
+
7. [Environmental Impact](#environmental-impact)
|
| 63 |
+
8. [Citation](#citation)
|
| 64 |
+
9. [Model Card Authors](#model-card-authors)
|
| 65 |
+
|
| 66 |
+
# TL;DR
|
| 67 |
+
|
| 68 |
+
If you already know T5, FLAN-T5 is just better at everything. For the same number of parameters, these models have been fine-tuned on more than 1000 additional tasks covering also more languages.
|
| 69 |
+
As mentioned in the first few lines of the abstract :
|
| 70 |
+
> Flan-PaLM 540B achieves state-of-the-art performance on several benchmarks, such as 75.2% on five-shot MMLU. We also publicly release Flan-T5 checkpoints,1 which achieve strong few-shot performance even compared to much larger models, such as PaLM 62B. Overall, instruction finetuning is a general method for improving the performance and usability of pretrained language models.
|
| 71 |
+
|
| 72 |
+
**Disclaimer**: Content from **this** model card has been written by the Hugging Face team, and parts of it were copy pasted from the [T5 model card](https://huggingface.co/t5-large).
|
| 73 |
+
|
| 74 |
+
# Model Details
|
| 75 |
+
|
| 76 |
+
## Model Description
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
- **Model type:** Language model
|
| 80 |
+
- **Language(s) (NLP):** English, Spanish, Japanese, Persian, Hindi, French, Chinese, Bengali, Gujarati, German, Telugu, Italian, Arabic, Polish, Tamil, Marathi, Malayalam, Oriya, Panjabi, Portuguese, Urdu, Galician, Hebrew, Korean, Catalan, Thai, Dutch, Indonesian, Vietnamese, Bulgarian, Filipino, Central Khmer, Lao, Turkish, Russian, Croatian, Swedish, Yoruba, Kurdish, Burmese, Malay, Czech, Finnish, Somali, Tagalog, Swahili, Sinhala, Kannada, Zhuang, Igbo, Xhosa, Romanian, Haitian, Estonian, Slovak, Lithuanian, Greek, Nepali, Assamese, Norwegian
|
| 81 |
+
- **License:** Apache 2.0
|
| 82 |
+
- **Related Models:** [All FLAN-T5 Checkpoints](https://huggingface.co/models?search=flan-t5)
|
| 83 |
+
- **Original Checkpoints:** [All Original FLAN-T5 Checkpoints](https://github.com/google-research/t5x/blob/main/docs/models.md#flan-t5-checkpoints)
|
| 84 |
+
- **Resources for more information:**
|
| 85 |
+
- [Research paper](https://arxiv.org/pdf/2210.11416.pdf)
|
| 86 |
+
- [GitHub Repo](https://github.com/google-research/t5x)
|
| 87 |
+
- [Hugging Face FLAN-T5 Docs (Similar to T5) ](https://huggingface.co/docs/transformers/model_doc/t5)
|
| 88 |
+
|
| 89 |
+
# Usage
|
| 90 |
+
|
| 91 |
+
Find below some example scripts on how to use the model in `transformers`:
|
| 92 |
+
|
| 93 |
+
## Using the Pytorch model
|
| 94 |
+
|
| 95 |
+
### Running the model on a CPU
|
| 96 |
+
|
| 97 |
+
<details>
|
| 98 |
+
<summary> Click to expand </summary>
|
| 99 |
+
|
| 100 |
+
```python
|
| 101 |
+
|
| 102 |
+
from transformers import T5Tokenizer, T5ForConditionalGeneration
|
| 103 |
+
|
| 104 |
+
tokenizer = T5Tokenizer.from_pretrained("google/flan-t5-large")
|
| 105 |
+
model = T5ForConditionalGeneration.from_pretrained("google/flan-t5-large")
|
| 106 |
+
|
| 107 |
+
input_text = "translate English to German: How old are you?"
|
| 108 |
+
input_ids = tokenizer(input_text, return_tensors="pt").input_ids
|
| 109 |
+
|
| 110 |
+
outputs = model.generate(input_ids)
|
| 111 |
+
print(tokenizer.decode(outputs[0]))
|
| 112 |
+
```
|
| 113 |
+
|
| 114 |
+
</details>
|
| 115 |
+
|
| 116 |
+
### Running the model on a GPU
|
| 117 |
+
|
| 118 |
+
<details>
|
| 119 |
+
<summary> Click to expand </summary>
|
| 120 |
+
|
| 121 |
+
```python
|
| 122 |
+
# pip install accelerate
|
| 123 |
+
from transformers import T5Tokenizer, T5ForConditionalGeneration
|
| 124 |
+
|
| 125 |
+
tokenizer = T5Tokenizer.from_pretrained("google/flan-t5-large")
|
| 126 |
+
model = T5ForConditionalGeneration.from_pretrained("google/flan-t5-large", device_map="auto")
|
| 127 |
+
|
| 128 |
+
input_text = "translate English to German: How old are you?"
|
| 129 |
+
input_ids = tokenizer(input_text, return_tensors="pt").input_ids.to("cuda")
|
| 130 |
+
|
| 131 |
+
outputs = model.generate(input_ids)
|
| 132 |
+
print(tokenizer.decode(outputs[0]))
|
| 133 |
+
```
|
| 134 |
+
|
| 135 |
+
</details>
|
| 136 |
+
|
| 137 |
+
### Running the model on a GPU using different precisions
|
| 138 |
+
|
| 139 |
+
#### FP16
|
| 140 |
+
|
| 141 |
+
<details>
|
| 142 |
+
<summary> Click to expand </summary>
|
| 143 |
+
|
| 144 |
+
```python
|
| 145 |
+
# pip install accelerate
|
| 146 |
+
import torch
|
| 147 |
+
from transformers import T5Tokenizer, T5ForConditionalGeneration
|
| 148 |
+
|
| 149 |
+
tokenizer = T5Tokenizer.from_pretrained("google/flan-t5-large")
|
| 150 |
+
model = T5ForConditionalGeneration.from_pretrained("google/flan-t5-large", device_map="auto", torch_dtype=torch.float16)
|
| 151 |
+
|
| 152 |
+
input_text = "translate English to German: How old are you?"
|
| 153 |
+
input_ids = tokenizer(input_text, return_tensors="pt").input_ids.to("cuda")
|
| 154 |
+
|
| 155 |
+
outputs = model.generate(input_ids)
|
| 156 |
+
print(tokenizer.decode(outputs[0]))
|
| 157 |
+
```
|
| 158 |
+
|
| 159 |
+
</details>
|
| 160 |
+
|
| 161 |
+
#### INT8
|
| 162 |
+
|
| 163 |
+
<details>
|
| 164 |
+
<summary> Click to expand </summary>
|
| 165 |
+
|
| 166 |
+
```python
|
| 167 |
+
# pip install bitsandbytes accelerate
|
| 168 |
+
from transformers import T5Tokenizer, T5ForConditionalGeneration
|
| 169 |
+
|
| 170 |
+
tokenizer = T5Tokenizer.from_pretrained("google/flan-t5-large")
|
| 171 |
+
model = T5ForConditionalGeneration.from_pretrained("google/flan-t5-large", device_map="auto", load_in_8bit=True)
|
| 172 |
+
|
| 173 |
+
input_text = "translate English to German: How old are you?"
|
| 174 |
+
input_ids = tokenizer(input_text, return_tensors="pt").input_ids.to("cuda")
|
| 175 |
+
|
| 176 |
+
outputs = model.generate(input_ids)
|
| 177 |
+
print(tokenizer.decode(outputs[0]))
|
| 178 |
+
```
|
| 179 |
+
|
| 180 |
+
</details>
|
| 181 |
+
|
| 182 |
+
# Uses
|
| 183 |
+
|
| 184 |
+
## Direct Use and Downstream Use
|
| 185 |
+
|
| 186 |
+
The authors write in [the original paper's model card](https://arxiv.org/pdf/2210.11416.pdf) that:
|
| 187 |
+
|
| 188 |
+
> The primary use is research on language models, including: research on zero-shot NLP tasks and in-context few-shot learning NLP tasks, such as reasoning, and question answering; advancing fairness and safety research, and understanding limitations of current large language models
|
| 189 |
+
|
| 190 |
+
See the [research paper](https://arxiv.org/pdf/2210.11416.pdf) for further details.
|
| 191 |
+
|
| 192 |
+
## Out-of-Scope Use
|
| 193 |
+
|
| 194 |
+
More information needed.
|
| 195 |
+
|
| 196 |
+
# Bias, Risks, and Limitations
|
| 197 |
+
|
| 198 |
+
The information below in this section are copied from the model's [official model card](https://arxiv.org/pdf/2210.11416.pdf):
|
| 199 |
+
|
| 200 |
+
> Language models, including Flan-T5, can potentially be used for language generation in a harmful way, according to Rae et al. (2021). Flan-T5 should not be used directly in any application, without a prior assessment of safety and fairness concerns specific to the application.
|
| 201 |
+
|
| 202 |
+
## Ethical considerations and risks
|
| 203 |
+
|
| 204 |
+
> Flan-T5 is fine-tuned on a large corpus of text data that was not filtered for explicit content or assessed for existing biases. As a result the model itself is potentially vulnerable to generating equivalently inappropriate content or replicating inherent biases in the underlying data.
|
| 205 |
+
|
| 206 |
+
## Known Limitations
|
| 207 |
+
|
| 208 |
+
> Flan-T5 has not been tested in real world applications.
|
| 209 |
+
|
| 210 |
+
## Sensitive Use:
|
| 211 |
+
|
| 212 |
+
> Flan-T5 should not be applied for any unacceptable use cases, e.g., generation of abusive speech.
|
| 213 |
+
|
| 214 |
+
# Training Details
|
| 215 |
+
|
| 216 |
+
## Training Data
|
| 217 |
+
|
| 218 |
+
The model was trained on a mixture of tasks, that includes the tasks described in the table below (from the original paper, figure 2):
|
| 219 |
+
|
| 220 |
+

|
| 221 |
+
|
| 222 |
+
|
| 223 |
+
## Training Procedure
|
| 224 |
+
|
| 225 |
+
According to the model card from the [original paper](https://arxiv.org/pdf/2210.11416.pdf):
|
| 226 |
+
|
| 227 |
+
> These models are based on pretrained T5 (Raffel et al., 2020) and fine-tuned with instructions for better zero-shot and few-shot performance. There is one fine-tuned Flan model per T5 model size.
|
| 228 |
+
|
| 229 |
+
The model has been trained on TPU v3 or TPU v4 pods, using [`t5x`](https://github.com/google-research/t5x) codebase together with [`jax`](https://github.com/google/jax).
|
| 230 |
+
|
| 231 |
+
|
| 232 |
+
# Evaluation
|
| 233 |
+
|
| 234 |
+
## Testing Data, Factors & Metrics
|
| 235 |
+
|
| 236 |
+
The authors evaluated the model on various tasks covering several languages (1836 in total). See the table below for some quantitative evaluation:
|
| 237 |
+

|
| 238 |
+
For full details, please check the [research paper](https://arxiv.org/pdf/2210.11416.pdf).
|
| 239 |
+
|
| 240 |
+
## Results
|
| 241 |
+
|
| 242 |
+
For full results for FLAN-T5-Large, see the [research paper](https://arxiv.org/pdf/2210.11416.pdf), Table 3.
|
| 243 |
+
|
| 244 |
+
# Environmental Impact
|
| 245 |
+
|
| 246 |
+
Carbon emissions can be estimated using the [Machine Learning Impact calculator](https://mlco2.github.io/impact#compute) presented in [Lacoste et al. (2019)](https://arxiv.org/abs/1910.09700).
|
| 247 |
+
|
| 248 |
+
- **Hardware Type:** Google Cloud TPU Pods - TPU v3 or TPU v4 | Number of chips β₯ 4.
|
| 249 |
+
- **Hours used:** More information needed
|
| 250 |
+
- **Cloud Provider:** GCP
|
| 251 |
+
- **Compute Region:** More information needed
|
| 252 |
+
- **Carbon Emitted:** More information needed
|
| 253 |
+
|
| 254 |
+
# Citation
|
| 255 |
+
|
| 256 |
+
**BibTeX:**
|
| 257 |
+
|
| 258 |
+
```bibtex
|
| 259 |
+
@misc{https://doi.org/10.48550/arxiv.2210.11416,
|
| 260 |
+
doi = {10.48550/ARXIV.2210.11416},
|
| 261 |
+
|
| 262 |
+
url = {https://arxiv.org/abs/2210.11416},
|
| 263 |
+
|
| 264 |
+
author = {Chung, Hyung Won and Hou, Le and Longpre, Shayne and Zoph, Barret and Tay, Yi and Fedus, William and Li, Eric and Wang, Xuezhi and Dehghani, Mostafa and Brahma, Siddhartha and Webson, Albert and Gu, Shixiang Shane and Dai, Zhuyun and Suzgun, Mirac and Chen, Xinyun and Chowdhery, Aakanksha and Narang, Sharan and Mishra, Gaurav and Yu, Adams and Zhao, Vincent and Huang, Yanping and Dai, Andrew and Yu, Hongkun and Petrov, Slav and Chi, Ed H. and Dean, Jeff and Devlin, Jacob and Roberts, Adam and Zhou, Denny and Le, Quoc V. and Wei, Jason},
|
| 265 |
+
|
| 266 |
+
keywords = {Machine Learning (cs.LG), Computation and Language (cs.CL), FOS: Computer and information sciences, FOS: Computer and information sciences},
|
| 267 |
+
|
| 268 |
+
title = {Scaling Instruction-Finetuned Language Models},
|
| 269 |
+
|
| 270 |
+
publisher = {arXiv},
|
| 271 |
+
|
| 272 |
+
year = {2022},
|
| 273 |
+
|
| 274 |
+
copyright = {Creative Commons Attribution 4.0 International}
|
| 275 |
+
}
|
| 276 |
+
```
|
Looped-DiT-B-32/text_encoder/config.json
ADDED
|
@@ -0,0 +1,28 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"architectures": [
|
| 3 |
+
"T5ForConditionalGeneration"
|
| 4 |
+
],
|
| 5 |
+
"d_ff": 2816,
|
| 6 |
+
"d_kv": 64,
|
| 7 |
+
"d_model": 1024,
|
| 8 |
+
"decoder_start_token_id": 0,
|
| 9 |
+
"dropout_rate": 0.1,
|
| 10 |
+
"eos_token_id": 1,
|
| 11 |
+
"feed_forward_proj": "gated-gelu",
|
| 12 |
+
"initializer_factor": 1.0,
|
| 13 |
+
"is_encoder_decoder": true,
|
| 14 |
+
"layer_norm_epsilon": 1e-06,
|
| 15 |
+
"model_type": "t5",
|
| 16 |
+
"n_positions": 512,
|
| 17 |
+
"num_decoder_layers": 24,
|
| 18 |
+
"num_heads": 16,
|
| 19 |
+
"num_layers": 24,
|
| 20 |
+
"output_past": true,
|
| 21 |
+
"pad_token_id": 0,
|
| 22 |
+
"relative_attention_max_distance": 128,
|
| 23 |
+
"relative_attention_num_buckets": 32,
|
| 24 |
+
"tie_word_embeddings": false,
|
| 25 |
+
"transformers_version": "4.23.1",
|
| 26 |
+
"use_cache": true,
|
| 27 |
+
"vocab_size": 32128
|
| 28 |
+
}
|
Looped-DiT-B-32/text_encoder/generation_config.json
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"_from_model_config": true,
|
| 3 |
+
"decoder_start_token_id": 0,
|
| 4 |
+
"eos_token_id": 1,
|
| 5 |
+
"pad_token_id": 0,
|
| 6 |
+
"transformers_version": "4.27.0.dev0"
|
| 7 |
+
}
|
Looped-DiT-B-32/text_encoder/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:149fe330e51e007c6ea9c6089f334ca82fcbe06b0df27e578e90934bcd327f73
|
| 3 |
+
size 3142856782
|
Looped-DiT-B-32/text_encoder/special_tokens_map.json
ADDED
|
@@ -0,0 +1,107 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"additional_special_tokens": [
|
| 3 |
+
"<extra_id_0>",
|
| 4 |
+
"<extra_id_1>",
|
| 5 |
+
"<extra_id_2>",
|
| 6 |
+
"<extra_id_3>",
|
| 7 |
+
"<extra_id_4>",
|
| 8 |
+
"<extra_id_5>",
|
| 9 |
+
"<extra_id_6>",
|
| 10 |
+
"<extra_id_7>",
|
| 11 |
+
"<extra_id_8>",
|
| 12 |
+
"<extra_id_9>",
|
| 13 |
+
"<extra_id_10>",
|
| 14 |
+
"<extra_id_11>",
|
| 15 |
+
"<extra_id_12>",
|
| 16 |
+
"<extra_id_13>",
|
| 17 |
+
"<extra_id_14>",
|
| 18 |
+
"<extra_id_15>",
|
| 19 |
+
"<extra_id_16>",
|
| 20 |
+
"<extra_id_17>",
|
| 21 |
+
"<extra_id_18>",
|
| 22 |
+
"<extra_id_19>",
|
| 23 |
+
"<extra_id_20>",
|
| 24 |
+
"<extra_id_21>",
|
| 25 |
+
"<extra_id_22>",
|
| 26 |
+
"<extra_id_23>",
|
| 27 |
+
"<extra_id_24>",
|
| 28 |
+
"<extra_id_25>",
|
| 29 |
+
"<extra_id_26>",
|
| 30 |
+
"<extra_id_27>",
|
| 31 |
+
"<extra_id_28>",
|
| 32 |
+
"<extra_id_29>",
|
| 33 |
+
"<extra_id_30>",
|
| 34 |
+
"<extra_id_31>",
|
| 35 |
+
"<extra_id_32>",
|
| 36 |
+
"<extra_id_33>",
|
| 37 |
+
"<extra_id_34>",
|
| 38 |
+
"<extra_id_35>",
|
| 39 |
+
"<extra_id_36>",
|
| 40 |
+
"<extra_id_37>",
|
| 41 |
+
"<extra_id_38>",
|
| 42 |
+
"<extra_id_39>",
|
| 43 |
+
"<extra_id_40>",
|
| 44 |
+
"<extra_id_41>",
|
| 45 |
+
"<extra_id_42>",
|
| 46 |
+
"<extra_id_43>",
|
| 47 |
+
"<extra_id_44>",
|
| 48 |
+
"<extra_id_45>",
|
| 49 |
+
"<extra_id_46>",
|
| 50 |
+
"<extra_id_47>",
|
| 51 |
+
"<extra_id_48>",
|
| 52 |
+
"<extra_id_49>",
|
| 53 |
+
"<extra_id_50>",
|
| 54 |
+
"<extra_id_51>",
|
| 55 |
+
"<extra_id_52>",
|
| 56 |
+
"<extra_id_53>",
|
| 57 |
+
"<extra_id_54>",
|
| 58 |
+
"<extra_id_55>",
|
| 59 |
+
"<extra_id_56>",
|
| 60 |
+
"<extra_id_57>",
|
| 61 |
+
"<extra_id_58>",
|
| 62 |
+
"<extra_id_59>",
|
| 63 |
+
"<extra_id_60>",
|
| 64 |
+
"<extra_id_61>",
|
| 65 |
+
"<extra_id_62>",
|
| 66 |
+
"<extra_id_63>",
|
| 67 |
+
"<extra_id_64>",
|
| 68 |
+
"<extra_id_65>",
|
| 69 |
+
"<extra_id_66>",
|
| 70 |
+
"<extra_id_67>",
|
| 71 |
+
"<extra_id_68>",
|
| 72 |
+
"<extra_id_69>",
|
| 73 |
+
"<extra_id_70>",
|
| 74 |
+
"<extra_id_71>",
|
| 75 |
+
"<extra_id_72>",
|
| 76 |
+
"<extra_id_73>",
|
| 77 |
+
"<extra_id_74>",
|
| 78 |
+
"<extra_id_75>",
|
| 79 |
+
"<extra_id_76>",
|
| 80 |
+
"<extra_id_77>",
|
| 81 |
+
"<extra_id_78>",
|
| 82 |
+
"<extra_id_79>",
|
| 83 |
+
"<extra_id_80>",
|
| 84 |
+
"<extra_id_81>",
|
| 85 |
+
"<extra_id_82>",
|
| 86 |
+
"<extra_id_83>",
|
| 87 |
+
"<extra_id_84>",
|
| 88 |
+
"<extra_id_85>",
|
| 89 |
+
"<extra_id_86>",
|
| 90 |
+
"<extra_id_87>",
|
| 91 |
+
"<extra_id_88>",
|
| 92 |
+
"<extra_id_89>",
|
| 93 |
+
"<extra_id_90>",
|
| 94 |
+
"<extra_id_91>",
|
| 95 |
+
"<extra_id_92>",
|
| 96 |
+
"<extra_id_93>",
|
| 97 |
+
"<extra_id_94>",
|
| 98 |
+
"<extra_id_95>",
|
| 99 |
+
"<extra_id_96>",
|
| 100 |
+
"<extra_id_97>",
|
| 101 |
+
"<extra_id_98>",
|
| 102 |
+
"<extra_id_99>"
|
| 103 |
+
],
|
| 104 |
+
"eos_token": "</s>",
|
| 105 |
+
"pad_token": "<pad>",
|
| 106 |
+
"unk_token": "<unk>"
|
| 107 |
+
}
|
Looped-DiT-B-32/text_encoder/spiece.model
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:89fa65b45c6c46d9ff3ecaf7a4eeff28d758a92d87bda5102dcb8141a0c051d3
|
| 3 |
+
size 859107
|
Looped-DiT-B-32/text_encoder/tokenizer.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
Looped-DiT-B-32/text_encoder/tokenizer_config.json
ADDED
|
@@ -0,0 +1,113 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"additional_special_tokens": [
|
| 3 |
+
"<extra_id_0>",
|
| 4 |
+
"<extra_id_1>",
|
| 5 |
+
"<extra_id_2>",
|
| 6 |
+
"<extra_id_3>",
|
| 7 |
+
"<extra_id_4>",
|
| 8 |
+
"<extra_id_5>",
|
| 9 |
+
"<extra_id_6>",
|
| 10 |
+
"<extra_id_7>",
|
| 11 |
+
"<extra_id_8>",
|
| 12 |
+
"<extra_id_9>",
|
| 13 |
+
"<extra_id_10>",
|
| 14 |
+
"<extra_id_11>",
|
| 15 |
+
"<extra_id_12>",
|
| 16 |
+
"<extra_id_13>",
|
| 17 |
+
"<extra_id_14>",
|
| 18 |
+
"<extra_id_15>",
|
| 19 |
+
"<extra_id_16>",
|
| 20 |
+
"<extra_id_17>",
|
| 21 |
+
"<extra_id_18>",
|
| 22 |
+
"<extra_id_19>",
|
| 23 |
+
"<extra_id_20>",
|
| 24 |
+
"<extra_id_21>",
|
| 25 |
+
"<extra_id_22>",
|
| 26 |
+
"<extra_id_23>",
|
| 27 |
+
"<extra_id_24>",
|
| 28 |
+
"<extra_id_25>",
|
| 29 |
+
"<extra_id_26>",
|
| 30 |
+
"<extra_id_27>",
|
| 31 |
+
"<extra_id_28>",
|
| 32 |
+
"<extra_id_29>",
|
| 33 |
+
"<extra_id_30>",
|
| 34 |
+
"<extra_id_31>",
|
| 35 |
+
"<extra_id_32>",
|
| 36 |
+
"<extra_id_33>",
|
| 37 |
+
"<extra_id_34>",
|
| 38 |
+
"<extra_id_35>",
|
| 39 |
+
"<extra_id_36>",
|
| 40 |
+
"<extra_id_37>",
|
| 41 |
+
"<extra_id_38>",
|
| 42 |
+
"<extra_id_39>",
|
| 43 |
+
"<extra_id_40>",
|
| 44 |
+
"<extra_id_41>",
|
| 45 |
+
"<extra_id_42>",
|
| 46 |
+
"<extra_id_43>",
|
| 47 |
+
"<extra_id_44>",
|
| 48 |
+
"<extra_id_45>",
|
| 49 |
+
"<extra_id_46>",
|
| 50 |
+
"<extra_id_47>",
|
| 51 |
+
"<extra_id_48>",
|
| 52 |
+
"<extra_id_49>",
|
| 53 |
+
"<extra_id_50>",
|
| 54 |
+
"<extra_id_51>",
|
| 55 |
+
"<extra_id_52>",
|
| 56 |
+
"<extra_id_53>",
|
| 57 |
+
"<extra_id_54>",
|
| 58 |
+
"<extra_id_55>",
|
| 59 |
+
"<extra_id_56>",
|
| 60 |
+
"<extra_id_57>",
|
| 61 |
+
"<extra_id_58>",
|
| 62 |
+
"<extra_id_59>",
|
| 63 |
+
"<extra_id_60>",
|
| 64 |
+
"<extra_id_61>",
|
| 65 |
+
"<extra_id_62>",
|
| 66 |
+
"<extra_id_63>",
|
| 67 |
+
"<extra_id_64>",
|
| 68 |
+
"<extra_id_65>",
|
| 69 |
+
"<extra_id_66>",
|
| 70 |
+
"<extra_id_67>",
|
| 71 |
+
"<extra_id_68>",
|
| 72 |
+
"<extra_id_69>",
|
| 73 |
+
"<extra_id_70>",
|
| 74 |
+
"<extra_id_71>",
|
| 75 |
+
"<extra_id_72>",
|
| 76 |
+
"<extra_id_73>",
|
| 77 |
+
"<extra_id_74>",
|
| 78 |
+
"<extra_id_75>",
|
| 79 |
+
"<extra_id_76>",
|
| 80 |
+
"<extra_id_77>",
|
| 81 |
+
"<extra_id_78>",
|
| 82 |
+
"<extra_id_79>",
|
| 83 |
+
"<extra_id_80>",
|
| 84 |
+
"<extra_id_81>",
|
| 85 |
+
"<extra_id_82>",
|
| 86 |
+
"<extra_id_83>",
|
| 87 |
+
"<extra_id_84>",
|
| 88 |
+
"<extra_id_85>",
|
| 89 |
+
"<extra_id_86>",
|
| 90 |
+
"<extra_id_87>",
|
| 91 |
+
"<extra_id_88>",
|
| 92 |
+
"<extra_id_89>",
|
| 93 |
+
"<extra_id_90>",
|
| 94 |
+
"<extra_id_91>",
|
| 95 |
+
"<extra_id_92>",
|
| 96 |
+
"<extra_id_93>",
|
| 97 |
+
"<extra_id_94>",
|
| 98 |
+
"<extra_id_95>",
|
| 99 |
+
"<extra_id_96>",
|
| 100 |
+
"<extra_id_97>",
|
| 101 |
+
"<extra_id_98>",
|
| 102 |
+
"<extra_id_99>"
|
| 103 |
+
],
|
| 104 |
+
"eos_token": "</s>",
|
| 105 |
+
"extra_ids": 100,
|
| 106 |
+
"model_max_length": 256,
|
| 107 |
+
"name_or_path": "google/t5-v1_1-large",
|
| 108 |
+
"pad_token": "<pad>",
|
| 109 |
+
"sp_model_kwargs": {},
|
| 110 |
+
"special_tokens_map_file": "/home/younes_huggingface_co/.cache/huggingface/hub/models--google--t5-v1_1-large/snapshots/314bc112b191ec17b625ba81438dc73d6c23659d/special_tokens_map.json",
|
| 111 |
+
"tokenizer_class": "T5Tokenizer",
|
| 112 |
+
"unk_token": "<unk>"
|
| 113 |
+
}
|
Looped-DiT-B-32/tokenizer/special_tokens_map.json
ADDED
|
@@ -0,0 +1,107 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"additional_special_tokens": [
|
| 3 |
+
"<extra_id_0>",
|
| 4 |
+
"<extra_id_1>",
|
| 5 |
+
"<extra_id_2>",
|
| 6 |
+
"<extra_id_3>",
|
| 7 |
+
"<extra_id_4>",
|
| 8 |
+
"<extra_id_5>",
|
| 9 |
+
"<extra_id_6>",
|
| 10 |
+
"<extra_id_7>",
|
| 11 |
+
"<extra_id_8>",
|
| 12 |
+
"<extra_id_9>",
|
| 13 |
+
"<extra_id_10>",
|
| 14 |
+
"<extra_id_11>",
|
| 15 |
+
"<extra_id_12>",
|
| 16 |
+
"<extra_id_13>",
|
| 17 |
+
"<extra_id_14>",
|
| 18 |
+
"<extra_id_15>",
|
| 19 |
+
"<extra_id_16>",
|
| 20 |
+
"<extra_id_17>",
|
| 21 |
+
"<extra_id_18>",
|
| 22 |
+
"<extra_id_19>",
|
| 23 |
+
"<extra_id_20>",
|
| 24 |
+
"<extra_id_21>",
|
| 25 |
+
"<extra_id_22>",
|
| 26 |
+
"<extra_id_23>",
|
| 27 |
+
"<extra_id_24>",
|
| 28 |
+
"<extra_id_25>",
|
| 29 |
+
"<extra_id_26>",
|
| 30 |
+
"<extra_id_27>",
|
| 31 |
+
"<extra_id_28>",
|
| 32 |
+
"<extra_id_29>",
|
| 33 |
+
"<extra_id_30>",
|
| 34 |
+
"<extra_id_31>",
|
| 35 |
+
"<extra_id_32>",
|
| 36 |
+
"<extra_id_33>",
|
| 37 |
+
"<extra_id_34>",
|
| 38 |
+
"<extra_id_35>",
|
| 39 |
+
"<extra_id_36>",
|
| 40 |
+
"<extra_id_37>",
|
| 41 |
+
"<extra_id_38>",
|
| 42 |
+
"<extra_id_39>",
|
| 43 |
+
"<extra_id_40>",
|
| 44 |
+
"<extra_id_41>",
|
| 45 |
+
"<extra_id_42>",
|
| 46 |
+
"<extra_id_43>",
|
| 47 |
+
"<extra_id_44>",
|
| 48 |
+
"<extra_id_45>",
|
| 49 |
+
"<extra_id_46>",
|
| 50 |
+
"<extra_id_47>",
|
| 51 |
+
"<extra_id_48>",
|
| 52 |
+
"<extra_id_49>",
|
| 53 |
+
"<extra_id_50>",
|
| 54 |
+
"<extra_id_51>",
|
| 55 |
+
"<extra_id_52>",
|
| 56 |
+
"<extra_id_53>",
|
| 57 |
+
"<extra_id_54>",
|
| 58 |
+
"<extra_id_55>",
|
| 59 |
+
"<extra_id_56>",
|
| 60 |
+
"<extra_id_57>",
|
| 61 |
+
"<extra_id_58>",
|
| 62 |
+
"<extra_id_59>",
|
| 63 |
+
"<extra_id_60>",
|
| 64 |
+
"<extra_id_61>",
|
| 65 |
+
"<extra_id_62>",
|
| 66 |
+
"<extra_id_63>",
|
| 67 |
+
"<extra_id_64>",
|
| 68 |
+
"<extra_id_65>",
|
| 69 |
+
"<extra_id_66>",
|
| 70 |
+
"<extra_id_67>",
|
| 71 |
+
"<extra_id_68>",
|
| 72 |
+
"<extra_id_69>",
|
| 73 |
+
"<extra_id_70>",
|
| 74 |
+
"<extra_id_71>",
|
| 75 |
+
"<extra_id_72>",
|
| 76 |
+
"<extra_id_73>",
|
| 77 |
+
"<extra_id_74>",
|
| 78 |
+
"<extra_id_75>",
|
| 79 |
+
"<extra_id_76>",
|
| 80 |
+
"<extra_id_77>",
|
| 81 |
+
"<extra_id_78>",
|
| 82 |
+
"<extra_id_79>",
|
| 83 |
+
"<extra_id_80>",
|
| 84 |
+
"<extra_id_81>",
|
| 85 |
+
"<extra_id_82>",
|
| 86 |
+
"<extra_id_83>",
|
| 87 |
+
"<extra_id_84>",
|
| 88 |
+
"<extra_id_85>",
|
| 89 |
+
"<extra_id_86>",
|
| 90 |
+
"<extra_id_87>",
|
| 91 |
+
"<extra_id_88>",
|
| 92 |
+
"<extra_id_89>",
|
| 93 |
+
"<extra_id_90>",
|
| 94 |
+
"<extra_id_91>",
|
| 95 |
+
"<extra_id_92>",
|
| 96 |
+
"<extra_id_93>",
|
| 97 |
+
"<extra_id_94>",
|
| 98 |
+
"<extra_id_95>",
|
| 99 |
+
"<extra_id_96>",
|
| 100 |
+
"<extra_id_97>",
|
| 101 |
+
"<extra_id_98>",
|
| 102 |
+
"<extra_id_99>"
|
| 103 |
+
],
|
| 104 |
+
"eos_token": "</s>",
|
| 105 |
+
"pad_token": "<pad>",
|
| 106 |
+
"unk_token": "<unk>"
|
| 107 |
+
}
|
Looped-DiT-B-32/tokenizer/spiece.model
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:89fa65b45c6c46d9ff3ecaf7a4eeff28d758a92d87bda5102dcb8141a0c051d3
|
| 3 |
+
size 859107
|
Looped-DiT-B-32/tokenizer/tokenizer.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
Looped-DiT-B-32/tokenizer/tokenizer_config.json
ADDED
|
@@ -0,0 +1,113 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"additional_special_tokens": [
|
| 3 |
+
"<extra_id_0>",
|
| 4 |
+
"<extra_id_1>",
|
| 5 |
+
"<extra_id_2>",
|
| 6 |
+
"<extra_id_3>",
|
| 7 |
+
"<extra_id_4>",
|
| 8 |
+
"<extra_id_5>",
|
| 9 |
+
"<extra_id_6>",
|
| 10 |
+
"<extra_id_7>",
|
| 11 |
+
"<extra_id_8>",
|
| 12 |
+
"<extra_id_9>",
|
| 13 |
+
"<extra_id_10>",
|
| 14 |
+
"<extra_id_11>",
|
| 15 |
+
"<extra_id_12>",
|
| 16 |
+
"<extra_id_13>",
|
| 17 |
+
"<extra_id_14>",
|
| 18 |
+
"<extra_id_15>",
|
| 19 |
+
"<extra_id_16>",
|
| 20 |
+
"<extra_id_17>",
|
| 21 |
+
"<extra_id_18>",
|
| 22 |
+
"<extra_id_19>",
|
| 23 |
+
"<extra_id_20>",
|
| 24 |
+
"<extra_id_21>",
|
| 25 |
+
"<extra_id_22>",
|
| 26 |
+
"<extra_id_23>",
|
| 27 |
+
"<extra_id_24>",
|
| 28 |
+
"<extra_id_25>",
|
| 29 |
+
"<extra_id_26>",
|
| 30 |
+
"<extra_id_27>",
|
| 31 |
+
"<extra_id_28>",
|
| 32 |
+
"<extra_id_29>",
|
| 33 |
+
"<extra_id_30>",
|
| 34 |
+
"<extra_id_31>",
|
| 35 |
+
"<extra_id_32>",
|
| 36 |
+
"<extra_id_33>",
|
| 37 |
+
"<extra_id_34>",
|
| 38 |
+
"<extra_id_35>",
|
| 39 |
+
"<extra_id_36>",
|
| 40 |
+
"<extra_id_37>",
|
| 41 |
+
"<extra_id_38>",
|
| 42 |
+
"<extra_id_39>",
|
| 43 |
+
"<extra_id_40>",
|
| 44 |
+
"<extra_id_41>",
|
| 45 |
+
"<extra_id_42>",
|
| 46 |
+
"<extra_id_43>",
|
| 47 |
+
"<extra_id_44>",
|
| 48 |
+
"<extra_id_45>",
|
| 49 |
+
"<extra_id_46>",
|
| 50 |
+
"<extra_id_47>",
|
| 51 |
+
"<extra_id_48>",
|
| 52 |
+
"<extra_id_49>",
|
| 53 |
+
"<extra_id_50>",
|
| 54 |
+
"<extra_id_51>",
|
| 55 |
+
"<extra_id_52>",
|
| 56 |
+
"<extra_id_53>",
|
| 57 |
+
"<extra_id_54>",
|
| 58 |
+
"<extra_id_55>",
|
| 59 |
+
"<extra_id_56>",
|
| 60 |
+
"<extra_id_57>",
|
| 61 |
+
"<extra_id_58>",
|
| 62 |
+
"<extra_id_59>",
|
| 63 |
+
"<extra_id_60>",
|
| 64 |
+
"<extra_id_61>",
|
| 65 |
+
"<extra_id_62>",
|
| 66 |
+
"<extra_id_63>",
|
| 67 |
+
"<extra_id_64>",
|
| 68 |
+
"<extra_id_65>",
|
| 69 |
+
"<extra_id_66>",
|
| 70 |
+
"<extra_id_67>",
|
| 71 |
+
"<extra_id_68>",
|
| 72 |
+
"<extra_id_69>",
|
| 73 |
+
"<extra_id_70>",
|
| 74 |
+
"<extra_id_71>",
|
| 75 |
+
"<extra_id_72>",
|
| 76 |
+
"<extra_id_73>",
|
| 77 |
+
"<extra_id_74>",
|
| 78 |
+
"<extra_id_75>",
|
| 79 |
+
"<extra_id_76>",
|
| 80 |
+
"<extra_id_77>",
|
| 81 |
+
"<extra_id_78>",
|
| 82 |
+
"<extra_id_79>",
|
| 83 |
+
"<extra_id_80>",
|
| 84 |
+
"<extra_id_81>",
|
| 85 |
+
"<extra_id_82>",
|
| 86 |
+
"<extra_id_83>",
|
| 87 |
+
"<extra_id_84>",
|
| 88 |
+
"<extra_id_85>",
|
| 89 |
+
"<extra_id_86>",
|
| 90 |
+
"<extra_id_87>",
|
| 91 |
+
"<extra_id_88>",
|
| 92 |
+
"<extra_id_89>",
|
| 93 |
+
"<extra_id_90>",
|
| 94 |
+
"<extra_id_91>",
|
| 95 |
+
"<extra_id_92>",
|
| 96 |
+
"<extra_id_93>",
|
| 97 |
+
"<extra_id_94>",
|
| 98 |
+
"<extra_id_95>",
|
| 99 |
+
"<extra_id_96>",
|
| 100 |
+
"<extra_id_97>",
|
| 101 |
+
"<extra_id_98>",
|
| 102 |
+
"<extra_id_99>"
|
| 103 |
+
],
|
| 104 |
+
"eos_token": "</s>",
|
| 105 |
+
"extra_ids": 100,
|
| 106 |
+
"model_max_length": 256,
|
| 107 |
+
"name_or_path": "google/t5-v1_1-large",
|
| 108 |
+
"pad_token": "<pad>",
|
| 109 |
+
"sp_model_kwargs": {},
|
| 110 |
+
"special_tokens_map_file": "/home/younes_huggingface_co/.cache/huggingface/hub/models--google--t5-v1_1-large/snapshots/314bc112b191ec17b625ba81438dc73d6c23659d/special_tokens_map.json",
|
| 111 |
+
"tokenizer_class": "T5Tokenizer",
|
| 112 |
+
"unk_token": "<unk>"
|
| 113 |
+
}
|
Looped-DiT-B-32/transformer/config.json
ADDED
|
@@ -0,0 +1,23 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"_class_name": "LoopedDiTTransformer2DModel",
|
| 3 |
+
"_diffusers_version": "0.39.0",
|
| 4 |
+
"head_dim": 64,
|
| 5 |
+
"hidden_size": 768,
|
| 6 |
+
"image_size": 512,
|
| 7 |
+
"in_channels": 3,
|
| 8 |
+
"loop_split": [
|
| 9 |
+
6,
|
| 10 |
+
5,
|
| 11 |
+
6
|
| 12 |
+
],
|
| 13 |
+
"mlp_ratio": 2.6667,
|
| 14 |
+
"num_heads": 12,
|
| 15 |
+
"num_loops": 4,
|
| 16 |
+
"patch_size": 32,
|
| 17 |
+
"pca_channels": 128,
|
| 18 |
+
"share_loop_weights": true,
|
| 19 |
+
"text_dim": 1024,
|
| 20 |
+
"text_preamble_depth": 2,
|
| 21 |
+
"use_attn_gate": false,
|
| 22 |
+
"use_xsa": true
|
| 23 |
+
}
|
Looped-DiT-B-32/transformer/diffusion_pytorch_model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:0044f19126f0ea046620695a2b0d8ae6f3ef77c5589f330f712bf0312df74dbb
|
| 3 |
+
size 1044176480
|
Looped-DiT-B-32/transformer/transformer_looped_dit.py
ADDED
|
@@ -0,0 +1,417 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Looped-DiT denoiser as a diffusers model.
|
| 2 |
+
|
| 3 |
+
The MiniT2I denoiser, a pixel-space variant of MMDiT (patchified image tokens
|
| 4 |
+
and T5 text tokens with modality-specific weights and joint attention, no
|
| 5 |
+
timestep conditioning), whose double-stream blocks are split into three stages:
|
| 6 |
+
|
| 7 |
+
pre-loop A blocks[:pre] run once
|
| 8 |
+
looped B the next `core` blocks run N times
|
| 9 |
+
post-loop C the last `post` blocks run once, followed by the head
|
| 10 |
+
|
| 11 |
+
h_0 = A(x), h_r = B(h_{r-1}), x0_hat(r) = C(h_r), r = 1..N
|
| 12 |
+
|
| 13 |
+
With shared weights (the default) B is one set of blocks reused N times, so the
|
| 14 |
+
loop adds depth but no weights. Any loop state h_r can be decoded through C:
|
| 15 |
+
deep supervision trains those intermediate exits, and inference can run with a
|
| 16 |
+
different loop depth. XSA and the attention gate act on the looped blocks only.
|
| 17 |
+
"""
|
| 18 |
+
|
| 19 |
+
from __future__ import annotations
|
| 20 |
+
|
| 21 |
+
import math
|
| 22 |
+
|
| 23 |
+
import torch
|
| 24 |
+
import torch.nn.functional as F
|
| 25 |
+
from diffusers.configuration_utils import ConfigMixin
|
| 26 |
+
from diffusers.models.modeling_utils import ModelMixin
|
| 27 |
+
from torch import nn
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
class RMSNorm(nn.Module):
|
| 31 |
+
def __init__(self, dim: int, eps: float = 1e-6):
|
| 32 |
+
super().__init__()
|
| 33 |
+
self.eps = eps
|
| 34 |
+
self.weight = nn.Parameter(torch.ones(dim))
|
| 35 |
+
|
| 36 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 37 |
+
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps) * self.weight
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
class SwiGLU(nn.Module):
|
| 41 |
+
def __init__(self, dim: int, hidden_dim: int):
|
| 42 |
+
super().__init__()
|
| 43 |
+
hidden_dim = math.ceil(hidden_dim / 8) * 8
|
| 44 |
+
self.w1 = nn.Linear(dim, hidden_dim, bias=False)
|
| 45 |
+
self.w3 = nn.Linear(dim, hidden_dim, bias=False)
|
| 46 |
+
self.w2 = nn.Linear(hidden_dim, dim, bias=False)
|
| 47 |
+
for layer in (self.w1, self.w3, self.w2):
|
| 48 |
+
nn.init.xavier_uniform_(layer.weight)
|
| 49 |
+
|
| 50 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 51 |
+
return self.w2(F.silu(self.w1(x)) * self.w3(x))
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
# ---------------------------------------------------------------------------
|
| 55 |
+
# Rotary position embeddings: 1D over text positions, 2D over the patch grid.
|
| 56 |
+
# ---------------------------------------------------------------------------
|
| 57 |
+
|
| 58 |
+
_ROPE_CACHE: dict[tuple, tuple[torch.Tensor, torch.Tensor]] = {}
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
def _autocast_state() -> tuple:
|
| 62 |
+
try:
|
| 63 |
+
enabled = torch.is_autocast_enabled("cuda")
|
| 64 |
+
return enabled, torch.get_autocast_dtype("cuda") if enabled else None
|
| 65 |
+
except TypeError: # torch < 2.4
|
| 66 |
+
enabled = torch.is_autocast_enabled()
|
| 67 |
+
return enabled, torch.get_autocast_gpu_dtype() if enabled else None
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
def _rope_tables(grid: int | None, n: int, d: int, device, dtype, theta: float = 10000.0):
|
| 71 |
+
"""cos/sin tables for 1D (grid=None) or 2D rotary embeddings.
|
| 72 |
+
|
| 73 |
+
Tables are built under whatever autocast state the caller runs in (under bf16
|
| 74 |
+
autocast the angle products are computed in bf16, which is what the models
|
| 75 |
+
were trained with), so that state is part of the cache key.
|
| 76 |
+
"""
|
| 77 |
+
key = (grid, n, d, str(device), dtype, _autocast_state())
|
| 78 |
+
if key not in _ROPE_CACHE:
|
| 79 |
+
if grid is None:
|
| 80 |
+
inv = 1.0 / (theta ** (torch.arange(0, d, 2, device=device, dtype=torch.float32) / d))
|
| 81 |
+
pos = torch.arange(n, device=device, dtype=torch.float32)
|
| 82 |
+
angles = torch.einsum("n,f->nf", pos, inv)
|
| 83 |
+
angles = torch.cat([angles, angles], dim=-1)
|
| 84 |
+
else:
|
| 85 |
+
half = d // 2
|
| 86 |
+
inv = 1.0 / (theta ** (torch.arange(0, half, 2, device=device, dtype=torch.float32) / half))
|
| 87 |
+
freqs = torch.einsum("n,f->nf", torch.arange(grid, device=device, dtype=torch.float32), inv)
|
| 88 |
+
f_h, f_w = torch.broadcast_tensors(freqs[:, None, :], freqs[None, :, :])
|
| 89 |
+
angles = torch.cat([f_h, f_w], dim=-1)
|
| 90 |
+
angles = torch.cat([angles, angles], dim=-1).reshape(n, d)
|
| 91 |
+
_ROPE_CACHE[key] = (angles.cos()[None, None].to(dtype), angles.sin()[None, None].to(dtype))
|
| 92 |
+
return _ROPE_CACHE[key]
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
def rotate_half(x: torch.Tensor) -> torch.Tensor:
|
| 96 |
+
x1, x2 = x.chunk(2, dim=-1)
|
| 97 |
+
return torch.cat([-x2, x1], dim=-1)
|
| 98 |
+
|
| 99 |
+
|
| 100 |
+
def apply_rope(x: torch.Tensor, grid: int | None = None) -> torch.Tensor:
|
| 101 |
+
"""x: [batch, heads, tokens, head_dim]."""
|
| 102 |
+
cos, sin = _rope_tables(grid, x.shape[2], x.shape[3], x.device, x.dtype)
|
| 103 |
+
return x * cos + rotate_half(x) * sin
|
| 104 |
+
|
| 105 |
+
|
| 106 |
+
def sincos_2d(dim: int, grid: int) -> torch.Tensor:
|
| 107 |
+
y, x = torch.meshgrid(torch.arange(grid), torch.arange(grid), indexing="ij")
|
| 108 |
+
omega = 1.0 / (10000 ** (torch.arange(dim // 4, dtype=torch.float32) / (dim // 4)))
|
| 109 |
+
out_y = torch.einsum("n,d->nd", y.flatten().float(), omega)
|
| 110 |
+
out_x = torch.einsum("n,d->nd", x.flatten().float(), omega)
|
| 111 |
+
return torch.cat([out_x.sin(), out_x.cos(), out_y.sin(), out_y.cos()], dim=1)
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
def exclusive_self_attention(out: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
|
| 115 |
+
"""XSA (Zhai, 2026: https://arxiv.org/abs/2603.09078):
|
| 116 |
+
remove from each token's attention output the component along that token's
|
| 117 |
+
own value vector, so attention only writes content from other tokens.
|
| 118 |
+
`out` and `v` are token-aligned, heads first."""
|
| 119 |
+
v_hat = F.normalize(v.float(), dim=-1)
|
| 120 |
+
out_f = out.float()
|
| 121 |
+
return (out_f - (out_f * v_hat).sum(dim=-1, keepdim=True) * v_hat).to(out.dtype)
|
| 122 |
+
|
| 123 |
+
|
| 124 |
+
# ---------------------------------------------------------------------------
|
| 125 |
+
# Blocks
|
| 126 |
+
# ---------------------------------------------------------------------------
|
| 127 |
+
|
| 128 |
+
|
| 129 |
+
class PatchEmbed(nn.Module):
|
| 130 |
+
"""Two-stage patch embedding: a low-rank patch projection, then a 1x1 conv."""
|
| 131 |
+
|
| 132 |
+
def __init__(self, patch_size: int, in_channels: int, hidden_size: int, bottleneck: int):
|
| 133 |
+
super().__init__()
|
| 134 |
+
self.proj1 = nn.Conv2d(in_channels, bottleneck, kernel_size=patch_size, stride=patch_size, bias=False)
|
| 135 |
+
self.proj2 = nn.Conv2d(bottleneck, hidden_size, kernel_size=1, bias=True)
|
| 136 |
+
nn.init.xavier_uniform_(self.proj1.weight)
|
| 137 |
+
nn.init.xavier_uniform_(self.proj2.weight)
|
| 138 |
+
nn.init.zeros_(self.proj2.bias)
|
| 139 |
+
|
| 140 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 141 |
+
return self.proj2(self.proj1(x)).flatten(2).transpose(1, 2)
|
| 142 |
+
|
| 143 |
+
|
| 144 |
+
class TextBlock(nn.Module):
|
| 145 |
+
"""Text-only transformer block that refines the T5 tokens before the joint blocks."""
|
| 146 |
+
|
| 147 |
+
def __init__(self, hidden_size: int, num_heads: int, head_dim: int, mlp_ratio: float):
|
| 148 |
+
super().__init__()
|
| 149 |
+
self.num_heads, self.head_dim = num_heads, head_dim
|
| 150 |
+
self.norm1 = RMSNorm(hidden_size)
|
| 151 |
+
self.norm2 = RMSNorm(hidden_size)
|
| 152 |
+
self.qkv = nn.Linear(hidden_size, num_heads * head_dim * 3)
|
| 153 |
+
self.proj = nn.Linear(num_heads * head_dim, hidden_size)
|
| 154 |
+
self.mlp = SwiGLU(hidden_size, int(hidden_size * mlp_ratio))
|
| 155 |
+
self.q_norm = RMSNorm(head_dim)
|
| 156 |
+
self.k_norm = RMSNorm(head_dim)
|
| 157 |
+
|
| 158 |
+
def forward(self, txt: torch.Tensor) -> torch.Tensor:
|
| 159 |
+
b, n, _ = txt.shape
|
| 160 |
+
q, k, v = self.qkv(self.norm1(txt)).view(b, n, 3, self.num_heads, self.head_dim).unbind(2)
|
| 161 |
+
q, k, v = (z.transpose(1, 2) for z in (self.q_norm(q), self.k_norm(k), v))
|
| 162 |
+
out = F.scaled_dot_product_attention(apply_rope(q), apply_rope(k), v, scale=self.head_dim**-0.5)
|
| 163 |
+
txt = txt + self.proj(out.transpose(1, 2).reshape(b, n, -1))
|
| 164 |
+
return txt + self.mlp(self.norm2(txt))
|
| 165 |
+
|
| 166 |
+
|
| 167 |
+
class DoubleStreamBlock(nn.Module):
|
| 168 |
+
"""MMDiT block: separate image/text weights, one joint attention over both.
|
| 169 |
+
|
| 170 |
+
`use_xsa` / `use_attn_gate` turn on self-modulating attention (set only for the
|
| 171 |
+
looped blocks). The attention gate (Qiu et al., 2026:
|
| 172 |
+
https://arxiv.org/abs/2505.06708) is head-wise:
|
| 173 |
+
y_i <- y_i * sigmoid(W_g u_i + b_g), with u_i the block's normed input.
|
| 174 |
+
|
| 175 |
+
`update_text=False` skips the text-stream update, for the last block, whose
|
| 176 |
+
text output is never read.
|
| 177 |
+
"""
|
| 178 |
+
|
| 179 |
+
def __init__(
|
| 180 |
+
self,
|
| 181 |
+
hidden_size: int,
|
| 182 |
+
num_heads: int,
|
| 183 |
+
head_dim: int,
|
| 184 |
+
mlp_ratio: float,
|
| 185 |
+
grid: int,
|
| 186 |
+
use_xsa: bool = False,
|
| 187 |
+
use_attn_gate: bool = False,
|
| 188 |
+
update_text: bool = True,
|
| 189 |
+
):
|
| 190 |
+
super().__init__()
|
| 191 |
+
self.num_heads, self.head_dim, self.grid = num_heads, head_dim, grid
|
| 192 |
+
self.use_xsa, self.use_attn_gate, self.update_text = use_xsa, use_attn_gate, update_text
|
| 193 |
+
inner = num_heads * head_dim
|
| 194 |
+
self.img_norm1 = RMSNorm(hidden_size)
|
| 195 |
+
self.img_norm2 = RMSNorm(hidden_size)
|
| 196 |
+
self.txt_norm1 = RMSNorm(hidden_size)
|
| 197 |
+
self.txt_norm2 = RMSNorm(hidden_size)
|
| 198 |
+
self.img_qkv = nn.Linear(hidden_size, inner * 3)
|
| 199 |
+
self.txt_qkv = nn.Linear(hidden_size, inner * 3)
|
| 200 |
+
self.q_norm = RMSNorm(head_dim)
|
| 201 |
+
self.k_norm = RMSNorm(head_dim)
|
| 202 |
+
self.img_proj = nn.Linear(inner, hidden_size)
|
| 203 |
+
self.txt_proj = nn.Linear(inner, hidden_size)
|
| 204 |
+
if use_attn_gate:
|
| 205 |
+
# Zero bias: the gates start half open on average.
|
| 206 |
+
self.img_gate = nn.Linear(hidden_size, num_heads)
|
| 207 |
+
self.txt_gate = nn.Linear(hidden_size, num_heads)
|
| 208 |
+
nn.init.zeros_(self.img_gate.bias)
|
| 209 |
+
nn.init.zeros_(self.txt_gate.bias)
|
| 210 |
+
self.img_mlp = SwiGLU(hidden_size, int(hidden_size * mlp_ratio))
|
| 211 |
+
self.txt_mlp = SwiGLU(hidden_size, int(hidden_size * mlp_ratio))
|
| 212 |
+
|
| 213 |
+
def forward(self, img: torch.Tensor, txt: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
| 214 |
+
b, li, _ = img.shape
|
| 215 |
+
lt = txt.shape[1]
|
| 216 |
+
img_n, txt_n = self.img_norm1(img), self.txt_norm1(txt)
|
| 217 |
+
qi, ki, vi = self.img_qkv(img_n).view(b, li, 3, self.num_heads, self.head_dim).unbind(2)
|
| 218 |
+
qt, kt, vt = self.txt_qkv(txt_n).view(b, lt, 3, self.num_heads, self.head_dim).unbind(2)
|
| 219 |
+
# Joint sequence [text; image], heads first.
|
| 220 |
+
q = torch.cat([qt, qi], dim=1).transpose(1, 2)
|
| 221 |
+
k = torch.cat([kt, ki], dim=1).transpose(1, 2)
|
| 222 |
+
v = torch.cat([vt, vi], dim=1).transpose(1, 2)
|
| 223 |
+
q = torch.cat([apply_rope(self.q_norm(q[:, :, :lt])), apply_rope(self.q_norm(q[:, :, lt:]), self.grid)], dim=2)
|
| 224 |
+
k = torch.cat([apply_rope(self.k_norm(k[:, :, :lt])), apply_rope(self.k_norm(k[:, :, lt:]), self.grid)], dim=2)
|
| 225 |
+
out = F.scaled_dot_product_attention(q, k, v, scale=self.head_dim**-0.5)
|
| 226 |
+
if self.use_xsa:
|
| 227 |
+
out = exclusive_self_attention(out, v)
|
| 228 |
+
out = out.transpose(1, 2) # [b, tokens, heads, head_dim]
|
| 229 |
+
out_t, out_i = out[:, :lt], out[:, lt:]
|
| 230 |
+
if self.use_attn_gate:
|
| 231 |
+
out_i = out_i * torch.sigmoid(self.img_gate(img_n)).unsqueeze(-1)
|
| 232 |
+
img = img + self.img_proj(out_i.reshape(b, li, -1))
|
| 233 |
+
img = img + self.img_mlp(self.img_norm2(img))
|
| 234 |
+
if self.update_text:
|
| 235 |
+
if self.use_attn_gate:
|
| 236 |
+
out_t = out_t * torch.sigmoid(self.txt_gate(txt_n)).unsqueeze(-1)
|
| 237 |
+
txt = txt + self.txt_proj(out_t.reshape(b, lt, -1))
|
| 238 |
+
txt = txt + self.txt_mlp(self.txt_norm2(txt))
|
| 239 |
+
return img, txt
|
| 240 |
+
|
| 241 |
+
|
| 242 |
+
# ---------------------------------------------------------------------------
|
| 243 |
+
# Looped MMDiT
|
| 244 |
+
# ---------------------------------------------------------------------------
|
| 245 |
+
|
| 246 |
+
|
| 247 |
+
class LoopedDiTTransformer2DModel(ModelMixin, ConfigMixin):
|
| 248 |
+
"""Predicts the clean image x0 from a noisy image and T5 text embeddings.
|
| 249 |
+
|
| 250 |
+
This is a diffusers [`ModelMixin`]: `save_pretrained` / `from_pretrained` round-trip the
|
| 251 |
+
architecture in `config.json`. `loop_split` is `(pre, core, post)` and `num_loops` is the
|
| 252 |
+
trained loop depth (`1` is the MiniT2I model without looping). With `share_loop_weights=False`
|
| 253 |
+
every pass gets its own copy of the core blocks: the compute-matched "deeper" baseline with
|
| 254 |
+
the same exits.
|
| 255 |
+
|
| 256 |
+
The forward pass does not take a timestep. Flow-matching time is applied by the pipeline when
|
| 257 |
+
it converts the x0 prediction into a velocity.
|
| 258 |
+
"""
|
| 259 |
+
|
| 260 |
+
def __init__(
|
| 261 |
+
self,
|
| 262 |
+
image_size: int = 512,
|
| 263 |
+
patch_size: int = 32,
|
| 264 |
+
in_channels: int = 3,
|
| 265 |
+
hidden_size: int = 768,
|
| 266 |
+
num_heads: int = 12,
|
| 267 |
+
head_dim: int = 64,
|
| 268 |
+
mlp_ratio: float = 2.6667,
|
| 269 |
+
pca_channels: int = 128,
|
| 270 |
+
text_dim: int = 1024,
|
| 271 |
+
text_preamble_depth: int = 2,
|
| 272 |
+
loop_split: tuple[int, int, int] | list[int] = (6, 5, 6),
|
| 273 |
+
num_loops: int = 4,
|
| 274 |
+
share_loop_weights: bool = True,
|
| 275 |
+
use_xsa: bool = False,
|
| 276 |
+
use_attn_gate: bool = False,
|
| 277 |
+
):
|
| 278 |
+
super().__init__()
|
| 279 |
+
loop_split = [int(n) for n in loop_split]
|
| 280 |
+
num_loops = int(num_loops)
|
| 281 |
+
share_loop_weights = bool(share_loop_weights)
|
| 282 |
+
use_xsa = bool(use_xsa)
|
| 283 |
+
use_attn_gate = bool(use_attn_gate)
|
| 284 |
+
if len(loop_split) != 3 or min(loop_split) < 1:
|
| 285 |
+
raise ValueError(f"loop_split must be three positive block counts (pre, core, post), got {loop_split}")
|
| 286 |
+
if num_loops < 1:
|
| 287 |
+
raise ValueError(f"num_loops must be >= 1, got {num_loops}")
|
| 288 |
+
# Lists (not tuples) so config.json stays valid JSON.
|
| 289 |
+
self.register_to_config(
|
| 290 |
+
image_size=image_size,
|
| 291 |
+
patch_size=patch_size,
|
| 292 |
+
in_channels=in_channels,
|
| 293 |
+
hidden_size=hidden_size,
|
| 294 |
+
num_heads=num_heads,
|
| 295 |
+
head_dim=head_dim,
|
| 296 |
+
mlp_ratio=mlp_ratio,
|
| 297 |
+
pca_channels=pca_channels,
|
| 298 |
+
text_dim=text_dim,
|
| 299 |
+
text_preamble_depth=text_preamble_depth,
|
| 300 |
+
loop_split=loop_split,
|
| 301 |
+
num_loops=num_loops,
|
| 302 |
+
share_loop_weights=share_loop_weights,
|
| 303 |
+
use_xsa=use_xsa,
|
| 304 |
+
use_attn_gate=use_attn_gate,
|
| 305 |
+
)
|
| 306 |
+
pre, core, post = loop_split
|
| 307 |
+
self.patch_size, self.in_channels = patch_size, in_channels
|
| 308 |
+
self.grid = image_size // patch_size
|
| 309 |
+
self.pre, self.core, self.post = pre, core, post
|
| 310 |
+
self.num_loops = num_loops
|
| 311 |
+
self.share_loop_weights = share_loop_weights
|
| 312 |
+
|
| 313 |
+
self.img_embed = PatchEmbed(patch_size, in_channels, hidden_size, pca_channels)
|
| 314 |
+
self.txt_embed = nn.Linear(text_dim, hidden_size, bias=False)
|
| 315 |
+
# Replaces the T5 embedding at padded prompt positions (and everywhere
|
| 316 |
+
# for the unconditional branch of classifier-free guidance).
|
| 317 |
+
self.mask_token = nn.Parameter(torch.zeros(1, 1, text_dim))
|
| 318 |
+
nn.init.normal_(self.mask_token, std=0.02)
|
| 319 |
+
# MiniT2I's timestep and pooled-text embedders. The model has no timestep
|
| 320 |
+
# conditioning and never uses them; they are kept (frozen, see below) so that
|
| 321 |
+
# the model and its checkpoints match MiniT2I and the paper.
|
| 322 |
+
self.t_embed = nn.ModuleDict(
|
| 323 |
+
{"mlp": nn.Sequential(nn.Linear(256, hidden_size), nn.SiLU(), nn.Linear(hidden_size, hidden_size))}
|
| 324 |
+
)
|
| 325 |
+
for layer in (self.t_embed.mlp[0], self.t_embed.mlp[2]):
|
| 326 |
+
nn.init.normal_(layer.weight, std=0.02)
|
| 327 |
+
nn.init.zeros_(layer.bias)
|
| 328 |
+
self.pooled_embed = nn.Linear(text_dim, hidden_size, bias=False)
|
| 329 |
+
self.register_buffer("pos_embed", sincos_2d(hidden_size, self.grid)[None], persistent=False)
|
| 330 |
+
self.txt_blocks = nn.ModuleList(
|
| 331 |
+
TextBlock(hidden_size, num_heads, head_dim, mlp_ratio) for _ in range(text_preamble_depth)
|
| 332 |
+
)
|
| 333 |
+
looped = core if self.share_loop_weights else core * self.num_loops
|
| 334 |
+
depth = pre + looped + post
|
| 335 |
+
self.blocks = nn.ModuleList(
|
| 336 |
+
DoubleStreamBlock(
|
| 337 |
+
hidden_size,
|
| 338 |
+
num_heads,
|
| 339 |
+
head_dim,
|
| 340 |
+
mlp_ratio,
|
| 341 |
+
self.grid,
|
| 342 |
+
use_xsa=use_xsa and pre <= i < pre + looped,
|
| 343 |
+
use_attn_gate=use_attn_gate and pre <= i < pre + looped,
|
| 344 |
+
update_text=i < depth - 1,
|
| 345 |
+
)
|
| 346 |
+
for i in range(depth)
|
| 347 |
+
)
|
| 348 |
+
self.final_norm = RMSNorm(hidden_size)
|
| 349 |
+
self.final = nn.Linear(hidden_size, patch_size * patch_size * in_channels)
|
| 350 |
+
nn.init.zeros_(self.final.weight)
|
| 351 |
+
nn.init.zeros_(self.final.bias)
|
| 352 |
+
# Frozen because they never receive a gradient: the unused embedders and the
|
| 353 |
+
# text-stream update of the last block (whose text output is never read).
|
| 354 |
+
last = self.blocks[-1]
|
| 355 |
+
for module in (self.t_embed, self.pooled_embed, last.txt_norm2, last.txt_proj, last.txt_mlp):
|
| 356 |
+
module.requires_grad_(False)
|
| 357 |
+
|
| 358 |
+
def unpatchify(self, x: torch.Tensor) -> torch.Tensor:
|
| 359 |
+
b, n, _ = x.shape
|
| 360 |
+
p, c, g = self.patch_size, self.in_channels, int(n**0.5)
|
| 361 |
+
x = x.view(b, g, g, p, p, c).permute(0, 5, 1, 3, 2, 4).contiguous()
|
| 362 |
+
return x.view(b, c, g * p, g * p)
|
| 363 |
+
|
| 364 |
+
def loop_blocks(self, r: int) -> nn.ModuleList:
|
| 365 |
+
"""Core blocks run on loop pass r (1-based)."""
|
| 366 |
+
start = self.pre if self.share_loop_weights else self.pre + (r - 1) * self.core
|
| 367 |
+
return self.blocks[start : start + self.core]
|
| 368 |
+
|
| 369 |
+
def decode(self, img: torch.Tensor, txt: torch.Tensor) -> torch.Tensor:
|
| 370 |
+
"""Post-loop blocks and output head: a loop state -> x0 prediction."""
|
| 371 |
+
for block in self.blocks[len(self.blocks) - self.post :]:
|
| 372 |
+
img, txt = block(img, txt)
|
| 373 |
+
return self.unpatchify(self.final(self.final_norm(img))).float()
|
| 374 |
+
|
| 375 |
+
def forward(
|
| 376 |
+
self,
|
| 377 |
+
x: torch.Tensor,
|
| 378 |
+
text: torch.Tensor,
|
| 379 |
+
text_mask: torch.Tensor,
|
| 380 |
+
num_loops: int | None = None,
|
| 381 |
+
exit_loops: tuple[int, ...] = (),
|
| 382 |
+
) -> torch.Tensor | tuple[torch.Tensor, dict[int, torch.Tensor]]:
|
| 383 |
+
"""x: noisy images [B, C, H, W]; text: T5 states [B, L, text_dim];
|
| 384 |
+
text_mask: [B, L], 1 for prompt tokens (all 0 = unconditional).
|
| 385 |
+
|
| 386 |
+
num_loops overrides the loop depth at inference. exit_loops lists
|
| 387 |
+
intermediate depths r < num_loops to decode as well; the call then
|
| 388 |
+
returns (final prediction, {r: prediction after r loops}).
|
| 389 |
+
"""
|
| 390 |
+
n = self.num_loops if num_loops is None else int(num_loops)
|
| 391 |
+
if n < 1 or (not self.share_loop_weights and n > self.num_loops):
|
| 392 |
+
raise ValueError(f"num_loops={n} is not available for this model (trained with {self.num_loops})")
|
| 393 |
+
exits = sorted({int(r) for r in exit_loops})
|
| 394 |
+
if any(not 1 <= r < n for r in exits):
|
| 395 |
+
raise ValueError(f"exit_loops must lie in [1, {n}), got {exits}")
|
| 396 |
+
|
| 397 |
+
text = torch.where(text_mask.to(torch.bool)[:, :, None], text, self.mask_token.to(text.dtype))
|
| 398 |
+
img = self.img_embed(x) + self.pos_embed.to(device=x.device, dtype=x.dtype)
|
| 399 |
+
txt = self.txt_embed(text)
|
| 400 |
+
for block in self.txt_blocks:
|
| 401 |
+
txt = block(txt)
|
| 402 |
+
for block in self.blocks[: self.pre]:
|
| 403 |
+
img, txt = block(img, txt)
|
| 404 |
+
states = {}
|
| 405 |
+
for r in range(1, n + 1):
|
| 406 |
+
for block in self.loop_blocks(r):
|
| 407 |
+
img, txt = block(img, txt)
|
| 408 |
+
if r in exits:
|
| 409 |
+
states[r] = (img, txt)
|
| 410 |
+
out = self.decode(img, txt)
|
| 411 |
+
if not exits:
|
| 412 |
+
return out
|
| 413 |
+
return out, {r: self.decode(*states[r]) for r in exits}
|
| 414 |
+
|
| 415 |
+
|
| 416 |
+
# Training imports this name. The diffusers class name is the one stored in checkpoints.
|
| 417 |
+
LoopedMMDiT = LoopedDiTTransformer2DModel
|
README.md
ADDED
|
@@ -0,0 +1,151 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
license: mit
|
| 3 |
+
library_name: diffusers
|
| 4 |
+
pipeline_tag: text-to-image
|
| 5 |
+
tags:
|
| 6 |
+
- diffusers
|
| 7 |
+
- looped-dit
|
| 8 |
+
- image-generation
|
| 9 |
+
- text-to-image
|
| 10 |
+
- flow-matching
|
| 11 |
+
- pixel-space
|
| 12 |
+
inference: true
|
| 13 |
+
widget:
|
| 14 |
+
- text: a red cube on top of a blue sphere
|
| 15 |
+
output:
|
| 16 |
+
url: Looped-DiT-B-16/demo.png
|
| 17 |
+
language:
|
| 18 |
+
- en
|
| 19 |
+
---
|
| 20 |
+
|
| 21 |
+
# BiliSakura/Looped-DiT-diffusers
|
| 22 |
+
|
| 23 |
+
Self-contained Looped-DiT text-to-image checkpoints for Hugging Face diffusers. Each variant folder ships its own pipeline code, component modules, bundled FLAN-T5-Large text encoder, and transformer weights.
|
| 24 |
+
|
| 25 |
+
## Available checkpoints
|
| 26 |
+
|
| 27 |
+
| Subfolder | Model | Params (denoiser + text encoder) | Patch | Loop depth | CFG |
|
| 28 |
+
| --- | --- | --- | ---: | ---: | ---: |
|
| 29 |
+
| [`Looped-DiT-B-32/`](Looped-DiT-B-32/) | Looped-DiT B/32 | 260M + 341M | 32 | 4 | 6.0 |
|
| 30 |
+
| [`Looped-DiT-B-16/`](Looped-DiT-B-16/) | Looped-DiT B/16 | 258M + 341M | 16 | 4 | 6.0 |
|
| 31 |
+
|
| 32 |
+
Benchmark scores (100 Euler steps, CFG 6.0, loop depth 4):
|
| 33 |
+
|
| 34 |
+
| Model | GenEval | DPG-Bench | PRISM | CoReBench | SpatialGenEval | TIIF-Short | Avg |
|
| 35 |
+
| --- | ---: | ---: | ---: | ---: | ---: | ---: | ---: |
|
| 36 |
+
| B/32 (290k) | 85.1 | 85.3 | 54.4 | 44.5 | 52.3 | 76.1 | 66.3 |
|
| 37 |
+
| B/16 (580k) | 87.4 | 87.0 | 67.0 | 53.5 | 54.6 | 79.7 | 71.5 |
|
| 38 |
+
|
| 39 |
+
## Repo layout
|
| 40 |
+
|
| 41 |
+
```text
|
| 42 |
+
BiliSakura/Looped-DiT-diffusers/
|
| 43 |
+
βββ README.md
|
| 44 |
+
βββ .gitattributes
|
| 45 |
+
βββ Looped-DiT-B-32/
|
| 46 |
+
β βββ pipeline.py
|
| 47 |
+
β βββ model_index.json
|
| 48 |
+
β βββ demo.png
|
| 49 |
+
β βββ scheduler/
|
| 50 |
+
β βββ text_encoder/
|
| 51 |
+
β βββ tokenizer/
|
| 52 |
+
β βββ transformer/
|
| 53 |
+
βββ Looped-DiT-B-16/
|
| 54 |
+
βββ ...
|
| 55 |
+
```
|
| 56 |
+
|
| 57 |
+
Each variant is self-contained: load with `custom_pipeline` pointing at that folderβs `pipeline.py` and `trust_remote_code=True`. Looped-DiT denoises directly in RGB pixel space (no VAE).
|
| 58 |
+
|
| 59 |
+
## Demo
|
| 60 |
+
|
| 61 |
+

|
| 62 |
+
|
| 63 |
+
Prompt: *"a red cube on top of a blue sphere."* β Looped-DiT B/16 at 512Γ512, 100 steps, `guidance_scale=6.0`, `num_loops=4`, `torch_dtype=bfloat16`, seed 42.
|
| 64 |
+
|
| 65 |
+

|
| 66 |
+
|
| 67 |
+
Same prompt and settings with Looped-DiT B/32.
|
| 68 |
+
|
| 69 |
+
## Load from Hugging Face
|
| 70 |
+
|
| 71 |
+
```python
|
| 72 |
+
import torch
|
| 73 |
+
from diffusers import DiffusionPipeline
|
| 74 |
+
|
| 75 |
+
pipe = DiffusionPipeline.from_pretrained(
|
| 76 |
+
"BiliSakura/Looped-DiT-diffusers",
|
| 77 |
+
subfolder="Looped-DiT-B-16",
|
| 78 |
+
custom_pipeline="pipeline.py",
|
| 79 |
+
trust_remote_code=True,
|
| 80 |
+
torch_dtype=torch.bfloat16,
|
| 81 |
+
).to("cuda")
|
| 82 |
+
|
| 83 |
+
generator = torch.Generator(device="cuda").manual_seed(42)
|
| 84 |
+
image = pipe(
|
| 85 |
+
"a red cube on top of a blue sphere",
|
| 86 |
+
num_inference_steps=100,
|
| 87 |
+
guidance_scale=6.0,
|
| 88 |
+
num_loops=4,
|
| 89 |
+
generator=generator,
|
| 90 |
+
).images[0]
|
| 91 |
+
image.save("demo.png")
|
| 92 |
+
```
|
| 93 |
+
|
| 94 |
+
For B/32, set `subfolder="Looped-DiT-B-32"`.
|
| 95 |
+
|
| 96 |
+
## Load from a local clone
|
| 97 |
+
|
| 98 |
+
```python
|
| 99 |
+
from pathlib import Path
|
| 100 |
+
import torch
|
| 101 |
+
from diffusers import DiffusionPipeline
|
| 102 |
+
|
| 103 |
+
model_dir = Path("./Looped-DiT-B-16").resolve()
|
| 104 |
+
pipe = DiffusionPipeline.from_pretrained(
|
| 105 |
+
str(model_dir),
|
| 106 |
+
local_files_only=True,
|
| 107 |
+
custom_pipeline=str(model_dir / "pipeline.py"),
|
| 108 |
+
trust_remote_code=True,
|
| 109 |
+
torch_dtype=torch.bfloat16,
|
| 110 |
+
).to("cuda")
|
| 111 |
+
|
| 112 |
+
generator = torch.Generator(device="cuda").manual_seed(42)
|
| 113 |
+
image = pipe(
|
| 114 |
+
"a red cube on top of a blue sphere",
|
| 115 |
+
num_inference_steps=100,
|
| 116 |
+
guidance_scale=6.0,
|
| 117 |
+
num_loops=4,
|
| 118 |
+
generator=generator,
|
| 119 |
+
).images[0]
|
| 120 |
+
image.save("demo.png")
|
| 121 |
+
```
|
| 122 |
+
|
| 123 |
+
Use `./Looped-DiT-B-32` instead of `./Looped-DiT-B-16` for the B/32 checkpoint.
|
| 124 |
+
|
| 125 |
+
## Recommended inference settings
|
| 126 |
+
|
| 127 |
+
| Variant | Resolution | Steps | CFG scale | `num_loops` | `torch_dtype` |
|
| 128 |
+
| --- | --- | ---: | ---: | ---: | --- |
|
| 129 |
+
| `Looped-DiT-B-32` | 512Γ512 | 100 | 6.0 | 4 (default) | `bfloat16` (full pipeline) |
|
| 130 |
+
| `Looped-DiT-B-16` | 512Γ512 | 100 | 6.0 | 4 (default) | `bfloat16` (full pipeline) |
|
| 131 |
+
|
| 132 |
+
Other loop depths work at inference when loop weights are shared (the default for released models).
|
| 133 |
+
|
| 134 |
+
## Interface notes
|
| 135 |
+
|
| 136 |
+
- Text conditioning uses bundled `google/flan-t5-large` (`T5EncoderModel` + `T5Tokenizer`) in **bfloat16**, the same dtype as the denoiser. Prompt length is the tokenizer `model_max_length` (256).
|
| 137 |
+
- `torch_dtype=torch.bfloat16` on `from_pretrained` sets both. Do not cast `pipe.text_encoder` back to float32.
|
| 138 |
+
- Set `custom_pipeline` to the variantβs `pipeline.py` (Hub: `"pipeline.py"` with `subfolder`; local: absolute path).
|
| 139 |
+
- Scheduler is `FlowMatchEulerDiscreteScheduler` with 1000 training timesteps and `shift=1.0`.
|
| 140 |
+
- `guidance_scale > 1.0` enables classifier-free guidance with an empty-string null prompt.
|
| 141 |
+
- Output resolution is fixed at 512Γ512.
|
| 142 |
+
|
| 143 |
+
## Links
|
| 144 |
+
|
| 145 |
+
- Upstream B/32 weights: [sensenova/Looped-DiT-B32](https://huggingface.co/sensenova/Looped-DiT-B32)
|
| 146 |
+
- Upstream B/16 weights: [sensenova/Looped-DiT-B16](https://huggingface.co/sensenova/Looped-DiT-B16)
|
| 147 |
+
- Backbone: [MiniT2I](https://github.com/PeppaKing8/minit2i-jax) Β· [BiliSakura/MiniT2I-diffusers](https://huggingface.co/BiliSakura/MiniT2I-diffusers)
|
| 148 |
+
|
| 149 |
+
## License
|
| 150 |
+
|
| 151 |
+
MIT (same as upstream Looped-DiT and MiniT2I).
|