BiliSakura commited on
Commit
020587f
Β·
verified Β·
1 Parent(s): 6c2ba58

Add files using upload-large-folder tool

Browse files
Files changed (40) hide show
  1. .gitattributes +37 -35
  2. Looped-DiT-B-16/demo.png +3 -0
  3. Looped-DiT-B-16/model_index.json +23 -0
  4. Looped-DiT-B-16/pipeline.py +763 -0
  5. Looped-DiT-B-16/scheduler/scheduler_config.json +18 -0
  6. Looped-DiT-B-16/text_encoder/README.md +276 -0
  7. Looped-DiT-B-16/text_encoder/config.json +28 -0
  8. Looped-DiT-B-16/text_encoder/generation_config.json +7 -0
  9. Looped-DiT-B-16/text_encoder/model.safetensors +3 -0
  10. Looped-DiT-B-16/text_encoder/special_tokens_map.json +107 -0
  11. Looped-DiT-B-16/text_encoder/spiece.model +3 -0
  12. Looped-DiT-B-16/text_encoder/tokenizer.json +0 -0
  13. Looped-DiT-B-16/text_encoder/tokenizer_config.json +113 -0
  14. Looped-DiT-B-16/tokenizer/special_tokens_map.json +107 -0
  15. Looped-DiT-B-16/tokenizer/spiece.model +3 -0
  16. Looped-DiT-B-16/tokenizer/tokenizer.json +0 -0
  17. Looped-DiT-B-16/tokenizer/tokenizer_config.json +113 -0
  18. Looped-DiT-B-16/transformer/config.json +23 -0
  19. Looped-DiT-B-16/transformer/diffusion_pytorch_model.safetensors +3 -0
  20. Looped-DiT-B-16/transformer/transformer_looped_dit.py +417 -0
  21. Looped-DiT-B-32/demo.png +3 -0
  22. Looped-DiT-B-32/model_index.json +23 -0
  23. Looped-DiT-B-32/pipeline.py +763 -0
  24. Looped-DiT-B-32/scheduler/scheduler_config.json +18 -0
  25. Looped-DiT-B-32/text_encoder/README.md +276 -0
  26. Looped-DiT-B-32/text_encoder/config.json +28 -0
  27. Looped-DiT-B-32/text_encoder/generation_config.json +7 -0
  28. Looped-DiT-B-32/text_encoder/model.safetensors +3 -0
  29. Looped-DiT-B-32/text_encoder/special_tokens_map.json +107 -0
  30. Looped-DiT-B-32/text_encoder/spiece.model +3 -0
  31. Looped-DiT-B-32/text_encoder/tokenizer.json +0 -0
  32. Looped-DiT-B-32/text_encoder/tokenizer_config.json +113 -0
  33. Looped-DiT-B-32/tokenizer/special_tokens_map.json +107 -0
  34. Looped-DiT-B-32/tokenizer/spiece.model +3 -0
  35. Looped-DiT-B-32/tokenizer/tokenizer.json +0 -0
  36. Looped-DiT-B-32/tokenizer/tokenizer_config.json +113 -0
  37. Looped-DiT-B-32/transformer/config.json +23 -0
  38. Looped-DiT-B-32/transformer/diffusion_pytorch_model.safetensors +3 -0
  39. Looped-DiT-B-32/transformer/transformer_looped_dit.py +417 -0
  40. 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

  • SHA256: 2761c015378004e748a13ca658c77650fd05d504c4af0a99bdbcf883164ee1d1
  • Pointer size: 131 Bytes
  • Size of remote file: 466 kB
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
+ ![table.png](https://s3.amazonaws.com/moonup/production/uploads/1666363265279-62441d1d9fdefb55a0b7d12c.png)
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
+ ![image.png](https://s3.amazonaws.com/moonup/production/uploads/1668072995230-62441d1d9fdefb55a0b7d12c.png)
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

  • SHA256: 596a20bc8304942aa0f25c247982256ef50f1118c27cb9a6600b36bef1bf1beb
  • Pointer size: 131 Bytes
  • Size of remote file: 468 kB
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
+ ![table.png](https://s3.amazonaws.com/moonup/production/uploads/1666363265279-62441d1d9fdefb55a0b7d12c.png)
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
+ ![image.png](https://s3.amazonaws.com/moonup/production/uploads/1668072995230-62441d1d9fdefb55a0b7d12c.png)
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
+ ![Looped-DiT-B-16 demo](Looped-DiT-B-16/demo.png)
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
+ ![Looped-DiT-B-32 demo](Looped-DiT-B-32/demo.png)
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).