Architectures · Vision & diffusion

Diffusion Transformer

DiT-XL/2 · predicting noise instead of the next token

compared
DiT-XL/2 vs GPT-2 small
layers
28 GPT-2 12
width
1,152 · 16 heads GPT-2 768 · 12
tokens
256 latent patches GPT-2 ≤ 1,024 text
output
noise per patch GPT-2 50,257 scores
params
675M · 118.6 GFLOPs GPT-2 124M

All steps

  1. GPT-2 vs DiT

    DiT keeps the Transformer block but changes what flows through it: noisy image patches instead of tokens, no mask, and the noise in each patch as the output. The timestep and class label are not tokens; they set each LayerNorm’s scale and shift and gate each sub-layer (adaLN-Zero). Click a label to jump.

    28 × 1,152 · 675M · 256 patches

  2. Adding noise

    Diffusion trains a model to undo noise. A clean image is mixed with Gaussian noise, a little at t = 1 and completely by t = 1,000; the model sees the noisy image and t and predicts the noise that was added. This is the exact mix on a toy image, with DiT’s schedule.

    xt = √ᾱt · x0 + √(1 − ᾱt) · ε

  3. Latent patches

    DiT works on a compressed image: a pretrained autoencoder turns 256 × 256 × 3 pixels into a 32 × 32 × 4 latent, and 2 × 2 patches of it make 256 tokens of 16 numbers. Smaller patches mean more tokens, more compute and better images.

    256² × 3 → 32² × 4 → 256 tokens

  4. adaLN-Zero

    In GPT-2, LayerNorm’s γ and β are fixed after training. In DiT a small network computes them, and a gate on each sub-layer’s output, from the timestep and the class, separately for every block. These are real DiT-XL/2 values; some sub-layers are almost switched off.

    x + α · f(LN(x) · (1 + γ) + β)

  5. Many passes per image

    Generating an image starts from pure noise and removes a little at a time: 250 steps, each a full pass through all 28 blocks, twice with classifier-free guidance. Here each step uses the true noise, as a perfect model would; the real model only estimates it.

    250 passes × 118.6 GFLOPs

Code

# the DiT block: the timestep and class set LayerNorm's scale and shift, and a gate
c = t_embedder(t) + y_embedder(y)                        # (1152,)
shift1, scale1, gate1, shift2, scale2, gate2 = adaLN_modulation(c).chunk(6)
x = x + gate1 * attn(layer_norm(x) * (1 + scale1) + shift1)
x = x + gate2 * mlp(layer_norm(x) * (1 + scale2) + shift2)
# training: add noise at a random t, predict it
x_t = alpha_bar[t].sqrt() * x0 + (1 - alpha_bar[t]).sqrt() * eps
loss = F.mse_loss(model(x_t, t, y), eps)
# sampling: 250 steps, each a full forward pass (twice with guidance)
for t in reversed(steps): x = p_sample(model, x, t, y)

Go deeper