Ir para o conteúdo

22. Diffusion Transformers

Diffusion Transformers (DiT)

Em 2023, Peebles & Xie1 demonstraram algo simples e impactante: a U-Net não é necessária para modelos de difusão. Substituindo-a por blocos Transformer puros, o modelo não só manteve a qualidade — ele passou a escalar previsivelmente com mais parâmetros e dados, exatamente como modelos de linguagem.

Hoje, toda geração de imagem e vídeo de ponta usa DiT:

Modelo Arquitetura Objetivo
FLUX.1 DiT (duplo-stream) Flow Matching
Stable Diffusion 3 MMDiT Flow Matching
Sora (OpenAI) Spacetime DiT Difusão
Movie Gen (Meta) DiT Flow Matching
CogVideoX DiT 3D Flow Matching

De U-Net para Transformer

A U-Net clássica usa convoluções com skip connections hierárquicas — boa para capturar detalhes locais, mas difícil de escalar. O DiT substitui tudo isso por blocos de atenção global.


Passo 1 — Patchify: Imagens como Sequências de Tokens

Assim como o ViT divide imagens em patches, o DiT opera no espaço latente (após o encoder VAE). Um latente de forma \(H \times W \times C\) é dividido em patches de tamanho \(p \times p\):

\[ \text{Número de tokens: } N = \frac{H}{p} \times \frac{W}{p} \]

Cada patch é linearizado e projetado para a dimensão \(d_{\text{model}}\) — tornando-se um "token visual".


Passo 2 — Bloco DiT com AdaLN

O DiT usa Adaptive Layer Normalization (AdaLN) para injetar a informação de timestep e classe/texto diretamente nos parâmetros de normalização:

\[ \text{AdaLN}(h, c) = \gamma(c) \cdot \frac{h - \mu}{\sigma} + \beta(c) \]

onde \(c = \text{MLP}(\text{emb}(t) + \text{emb}(\text{classe}))\) é o vetor de condicionamento.

Os parâmetros \(\gamma\) e \(\beta\) são preditos — não aprendidos estaticamente — tornando a normalização sensível ao passo de difusão e ao prompt.


Passo 3 — MMDiT: Atenção Bidirecional Multi-Modal

MMDiT (SD3, FLUX) vai além do condicionamento por cross-attention. Texto e imagem participam da mesma operação de atenção:

\[ [Q_{img} \| Q_{txt}] \cdot [K_{img} \| K_{txt}]^\top \]

Os tokens de imagem veem os tokens de texto e vice-versa — condicionamento muito mais rico do que injetar texto apenas via cross-attention.

O FLUX usa um design de "duplo stream": pesos separados para imagem e texto nos blocos Q/K/V/FFN, mas atenção compartilhada:

Stream img:  x_img → W_q^img·x  ─┐
                                   ├─→ concat → Atenção(Q,K,V) → separar
Stream txt:  x_txt → W_q^txt·x  ─┘

Visualização: Processo Completo de Geração


O que o DiT de fato comprou

A contribuição real do artigo do DiT não é "um Transformer também funciona aqui". É que trocar o backbone deu à geração de imagens a propriedade que a modelagem de linguagem já tinha: uma lei de escala com a qual se pode planejar1. Peebles e Xie mostraram o FID caindo de forma suave e previsível com o compute de treino, em quatro tamanhos de modelo e três tamanhos de patch, sem sinal do platô em que as U-Nets batem.

Foi isso que tornou a geração seguinte possível. Não se justifica um modelo de imagem de 12B de parâmetros sem uma curva dizendo o que 12B compram.

Modelo Parâmetros Tokens \(d\) Blocos
DiT-XL/2 (2023) 675M 256 1152 28
SD3 medium (2024) 2B 1024 1536 24
FLUX.1-dev (2024) 12B 4096 3072 57

As contagens de tokens são para a resolução de treino de cada modelo e são a grandeza que importa: a profundidade do SD3 acompanha a largura por \(d = 64 \times \text{blocos}\), então a configuração de 8B é o mesmo desenho com 38 blocos e 2432 canais.

Três propriedades vieram junto com a troca e cada uma importa mais que o número de FID:

  • Uma arquitetura para tudo. O mesmo bloco, os mesmos kernels, a mesma maquinaria de treino distribuído de um LLM. Paralelismo de sequência, FlashAttention, checkpointing de ativações, FSDP — tudo transfere sem mudanças.
  • Modalidade é só mais tokens. Acrescentar quadros de vídeo, áudio ou uma segunda imagem é concatenação, não uma arquitetura nova. É toda a razão pela qual modelos de vídeo são DiTs.
  • Resolução é um comprimento de sequência. Nenhum compromisso arquitetural com 512 ou 1024; você muda \(N\). O custo é o \(O(N^2)\) do capítulo 13 e é por isso que DiTs de alta resolução se apoiam em latentes de 16 canais e patches maiores em vez de mais pixels.

A U-Net não morreu e o DiT não é de graça

Em escala pequena e com pouco compute, uma U-Net bem ajustada ainda vence — o prior convolucional vale mais que atenção global quando não se pode pagar os dados (capítulo 10). O DiT vence em escala, que é exatamente o trade-off que este curso não para de encontrar. E o \(O(N^2)\) é real: uma passagem do FLUX com 4096 tokens em 1024² é dominada pela atenção e por isso todo DiT de alta resolução entrega alguma combinação de patches maiores, latentes mais ricos e kernels de atenção eficientes.

Para onde o desenho está indo

Desenvolvimento O que muda
MMDiT (SD3) Pesos separados para tokens de texto e de imagem, com atenção conjunta bidirecional. O texto deixa de ser entrada lateral somente-leitura e é por isso que o SD3 consegue renderizar palavras legíveis.
REPA5 Alinhar features intermediárias do DiT com um encoder auto-supervisionado congelado (DINOv2). Treina o mesmo modelo até uma ordem de grandeza mais rápido — o maior ganho barato publicado recentemente.
DiT de vídeo Dividir em patches o espaço e o tempo. Modelos classe Sora são isto: um DiT sobre patches latentes espaço-temporais, com o mesmo bloco.
Profundidade latente em vez de resolução VAEs de 16 canais em vez de mais tokens, para manter \(N\) acessível elevando a fidelidade.

Implementação Simplificada

import torch
import torch.nn as nn

class AdaLN(nn.Module):
    def __init__(self, d_model, d_cond):
        super().__init__()
        self.norm = nn.LayerNorm(d_model, elementwise_affine=False)
        self.proj = nn.Linear(d_cond, 2 * d_model)  # → γ, β

    def forward(self, x, c):
        gamma, beta = self.proj(c).chunk(2, dim=-1)
        return (1 + gamma.unsqueeze(1)) * self.norm(x) + beta.unsqueeze(1)

class DiTBlock(nn.Module):
    def __init__(self, d_model, n_heads, d_ff, d_cond):
        super().__init__()
        self.adaln1 = AdaLN(d_model, d_cond)
        self.attn = nn.MultiheadAttention(d_model, n_heads, batch_first=True)
        self.adaln2 = AdaLN(d_model, d_cond)
        self.ff = nn.Sequential(nn.Linear(d_model, d_ff), nn.GELU(), nn.Linear(d_ff, d_model))

    def forward(self, x, c):
        h = self.adaln1(x, c)
        x = x + self.attn(h, h, h)[0]      # self-attention
        x = x + self.ff(self.adaln2(x, c)) # FFN
        return x

class DiT(nn.Module):
    def __init__(self, in_channels, patch_size, d_model, n_heads, d_ff, n_layers, d_cond):
        super().__init__()
        self.patch_size = patch_size
        p = patch_size
        self.patchify = nn.Conv2d(in_channels, d_model, p, stride=p)
        self.blocks = nn.ModuleList([DiTBlock(d_model, n_heads, d_ff, d_cond) for _ in range(n_layers)])
        self.norm_out = nn.LayerNorm(d_model)
        self.depatchify = nn.Linear(d_model, p*p*in_channels)

    def forward(self, x, t_emb, cond):
        # x: (B, C, H, W) latente ruidoso
        B, C, H, W = x.shape
        tokens = self.patchify(x)                       # (B, d, H/p, W/p)
        tokens = tokens.flatten(2).transpose(1, 2)      # (B, N, d)
        c = t_emb + cond                                # combinar condicionamento
        for block in self.blocks:
            tokens = block(tokens, c)
        tokens = self.norm_out(tokens)
        patches = self.depatchify(tokens)               # (B, N, p*p*C)
        # reformatar para (B, C, H, W)
        p = self.patch_size
        patches = patches.view(B, H//p, W//p, p, p, C).permute(0,5,1,3,2,4).reshape(B,C,H,W)
        return patches  # campo de velocidade predito v_θ(x_t, t)

Pontos principais

  1. Um DiT é um ViT operando sobre patches latentes ruidosos. Divida em patches, some posição, rode \(L\) blocos Transformer e desfaça os patches para uma previsão de velocidade ou de ruído.
  2. O AdaLN-Zero injeta condicionamento modulando a normalização em vez de concatenar um token e o portão inicializado em zero faz cada bloco começar como identidade — o mesmo truque de adição segura do ControlNet e do LoRA.
  3. A contribuição é uma lei de escala para geração de imagens. FID suave e previsível contra compute foi o que justificou modelos de imagem de 12B de parâmetros.
  4. O MMDiT faz do texto um participante pleno, com pesos próprios e atenção conjunta. Foi daí que veio texto legível em imagens geradas.
  5. Uma arquitetura agora cobre linguagem, visão, geração de imagens e vídeo e todo investimento em infraestrutura transfere entre elas.
  6. O DiT vence em escala; uma U-Net ainda vence quando dados e compute são escassos. E o \(O(N^2)\) em tokens é a restrição que molda todo projeto de alta resolução.



  1. Peebles, W., & Xie, S. (2023). Scalable Diffusion Models with Transformers. ICCV 2023. ↩↩

  2. Esser, P. et al. (2024). Scaling Rectified Flow Transformers for High-Resolution Image Synthesis (SD3). ↩

  3. Black Forest Labs. (2024). FLUX.1. ↩

  4. Dosovitskiy, A. et al. (2021). An Image is Worth 16×16 Words: Transformers for Image Recognition at Scale. ↩

  5. Yu, S., et al. (2025). Representation Alignment for Generation: Training Diffusion Transformers Is Easier Than You Think — ICLR. REPA: alinhe features intermediárias com o DINOv2 e treine uma ordem de grandeza mais rápido. ↩