深入理解 Diffusion Transformer (DiT):扩散模型的 Transformer 范式

4834 字
24 分钟
深入理解 Diffusion Transformer (DiT):扩散模型的 Transformer 范式

1. DiT 的核心思想#

1.1 一个让人惊讶的事实#

2020-2022 年间,扩散模型在图像生成上”大杀四方”——DALL·E 2、Imagen、Stable Diffusion 接连发布,看起来已稳坐王座。

但所有这些模型都有一个共同点:

它们的核心去噪网络都是 UNet——一个专为图像设计的卷积架构。

这就引出一个自然而大胆的问题:

如果 Transformer 已经”统治”NLP、统一了视觉理解(ViT)、推翻了卷积回归——它能不能”接管”扩散模型的核心网络?

2022 年底,UC Berkeley & NYU 的 William Peebles 与 Saining Xie 给出了回答:

Diffusion Transformer (DiT)——用纯 Transformer 替代 UNet,在 ImageNet 256×256 上以 FID 1.51 击败了 LDM、ADM 等基于 UNet 的所有扩散模型,并展现出极佳的 Gflops-vs-FID 缩放曲线。

1.2 DiT 的核心思想#

一句话概括:

把扩散模型的噪声预测网络从卷积 UNet 替换为 Vision Transformer (ViT),并通过自适应层归一化 (AdaLN-Zero) 高效注入时间步与类别条件。

┌────────────────────────────────────────────────────────────────┐
│ DiT 整体架构示意 │
├────────────────────────────────────────────────────────────────┤
│ │
│ 噪声图像 x_T │
│ │ │
│ ▼ │
│ Patchify: (H/P)×(W/P) × D │
│ │ │
│ ▼ │
│ ┌──────────────────────────────┐ │
│ │ + Position Embedding │ 位置编码 │
│ └──────────────────────────────┘ │
│ │ │
│ ▼ │
│ ┌──────────────────────────────────────────────────┐ │
│ │ N × DiT Block │ │
│ │ ┌──────────────────────────┐ │ │
│ │ │ AdaLN-Zero (调制) │ ← (t, y) │ │
│ │ ├──────────────────────────┤ │ │
│ │ │ Multi-Head Self-Attention │ │ │
│ │ ├──────────────────────────┤ │ │
│ │ │ AdaLN-Zero │ ← (t, y) │ │
│ │ ├──────────────────────────┤ │ │
│ │ │ MLP │ │ │
│ │ └──────────────────────────┘ │ │
│ └──────────────────────────────────────────────────┘ │
│ │ │
│ ▼ │
│ Linear Decode: → (H/P)×(W/P) 噪声预测 │
│ │
└────────────────────────────────────────────────────────────────┘

1.3 一个生活化的比喻#

把 DiT 与 UNet 的差别想象成两个画家

  • UNet 画家:用固定大小的”画笔”(卷积核)一寸一寸画——卷积核小,效率高,但必须从局部到全局逐步建立结构
  • DiT 画家:先看一眼整张画(把所有 patch 全局关联起来),第一笔就能把不同区域关联在一起——更”聪明”,但需要更多计算和更多数据

1.4 DiT 在生成式 AI 历史上的位置#

2015-2019: GAN 主导图像生成
2020: DDPM 提出 (Ho et al.) - 重新让扩散模型回到舞台
│ 使用 UNet + 加性噪声预测
2021: LDM / Stable Diffusion
│ 把扩散搬到潜空间 (latent space)
│ 仍用 UNet
2022: DiT (Peebles & Xie, 2022.12)
│ **首次: 用纯 Transformer 取代 UNet**
│ 首次发现: 缩放律 (Gflops vs FID) 极优
2023: SD3 (MM-DiT)、PixArt-α、SDXL-Turbo
│ DiT 思想扩展到文本到图像
2024: Sora、Wan2.1、Open Sora
│ DiT 扩展到视频生成
│ Transformer 彻底统一了视觉生成
2025+: 视觉大模型的基础架构

2. 为什么可以替换 UNet?#

2.1 UNet 的”假设”#

UNet 在扩散模型中之所以好用,是因为它有恰到好处的归纳偏置 (Inductive Bias)

UNet 的偏置对扩散的贡献
局部性像素相邻关系强,卷积核捕获短程结构
平移不变性同一物体出现在图像任意位置,模型处理方式一致
金字塔结构多尺度感受野 + skip connection 保留细节
参数共享推理不同时间步很高效

这些偏置在小数据时代非常珍贵——它们让模型”少走弯路”。

2.2 DiT 的反问#

DiT 的论文提出了一个深刻的观察:

数据规模模型规模足够大时,归纳偏置的重要性会下降——Transformer 可以自己”学习到”这些偏置。

而在扩散生成任务上:

  • 已经有海量图文配对数据 (LAION 等)
  • 算力上不封顶 (DiT-XL 用了 8 卡 A100 训练若干月)
  • 高分辨率图像对”全局理解”要求高——正是 Transformer 的强项

2.3 UNet vs DiT 对比#

维度UNetDiT
核心算子2D 卷积 (ResNet block)Self-Attention + MLP
归纳偏置强 (局部、平移不变)弱 (全局,由数据学习)
信息传递跳层连接 + 多尺度全局自注意力
时间/类别条件拼接或 AdaGNAdaLN-Zero (更优)
数据效率小数据更友好大数据更优
缩放行为很快饱和强大、可预测
下限 vs 上限上限相对受限上限更高

3. DiT 架构详解#

3.1 整体流程图#

输入: 噪声图像 x_t (B, C, H, W) + 时间步 t + 类别标签 y
┌────────────────────────────────────────┘
▼ ▼
┌─────────┐ ┌──────────┐
│ Patchify│ │ t Embed │ + y Embed
└────┬────┘ └────┬─────┘
│ (B, S, D) │ (B, D_cond)
▼ ▼
┌──────────────────────────────┐ ┌────────────┐
│ + Pos. Embedding │ │ MLP │
└──────────────────────────────┘ └─────┬──────┘
│ │
▼ │
┌─────────────────────────────────────────────┘
┌──────────────────────────┐
│ N × DiT Block │ 对每个 block:
│ ├ AdaLN-Zero (调制) │ c ← (t+y) Embedding
│ ├ Self-Attention │ x ← x + SelfAttn(AdaLN(c)调制后的 x)
│ └ MLP │ x ← x + MLP(AdaLN(c)调制后的 x)
└──────────────────────────┘
Linear Decoder: 还原为噪声预测 (B, C, H, W)

3.2 Patchify (Patch 嵌入)#

DiT 的第一步是把 H×WH \times W 的图像转成 token 序列(与 ViT 类似)。设 patch 大小为 pp,则序列长度:

S=(Hp)×(Wp)S = \left(\frac{H}{p}\right) \times \left(\frac{W}{p}\right)
输入尺寸Patch序列长度 SS
32×32464
64×644256
256×25681024
256×25616256
import torch
import torch.nn as nn
class PatchEmbed(nn.Module):
"""将图像切分为 patch (DiT 版)。"""
def __init__(self, img_size=32, patch_size=4, in_chans=3, embed_dim=768):
super().__init__()
self.img_size = img_size
self.patch_size = patch_size
self.num_patches = (img_size // patch_size) ** 2
# 用 Conv2d 一气呵成: 切块 + 线性投影
self.proj = nn.Conv2d(
in_chans, embed_dim,
kernel_size=patch_size, stride=patch_size,
)
def forward(self, x):
"""
x: (B, C, H, W)
return: (B, S, D)
"""
x = self.proj(x) # (B, D, H/p, W/p)
x = x.flatten(2) # (B, D, S)
x = x.transpose(1, 2) # (B, S, D)
return x

3.3 时间步与类别嵌入#

class TimestepEmbedder(nn.Module):
"""时间步 t → 向量 (用正弦/余弦位置编码 + MLP)。"""
def __init__(self, hidden_size, frequency_embedding_size=256):
super().__init__()
self.mlp = nn.Sequential(
nn.Linear(frequency_embedding_size, hidden_size),
nn.SiLU(),
nn.Linear(hidden_size, hidden_size),
)
self.frequency_embedding_size = frequency_embedding_size
@staticmethod
def timestep_embedding(t, dim, max_period=10000):
"""Sinusoidal position embedding for t."""
half = dim // 2
freqs = torch.exp(
-math.log(max_period)
* torch.arange(half, dtype=torch.float32) / half
).to(t.device)
args = t[:, None].float() * freqs[None]
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
return embedding
def forward(self, t):
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
return self.mlp(t_freq)
class LabelEmbedder(nn.Module):
"""类别标签 y → 向量。DiT 默认使用 class-conditional 的 CFG。"""
def __init__(self, num_classes, hidden_size):
super().__init__()
# +1 是为"无条件" (CFG 时使用)
self.embedding = nn.Embedding(num_classes + 1, hidden_size)
def forward(self, labels):
return self.embedding(labels)

把时间步和类别 embedding 相加,得到最终条件向量 cc

c = t_embedder(t) + y_embedder(y)

3.4 自注意力块 (Self-Attention)#

与 ViT 几乎完全相同——QKV 投影 + 多头 softmax-attention + 输出投影。

class Attention(nn.Module):
"""Multi-Head Self-Attention (pre-LN)。"""
def __init__(self, dim, num_heads=8, qkv_bias=False):
super().__init__()
assert dim % num_heads == 0
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.scale = self.head_dim ** -0.5
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
self.proj = nn.Linear(dim, dim)
def forward(self, x):
B, N, C = x.shape
qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim)
qkv = qkv.permute(2, 0, 3, 1, 4) # (3, B, h, N, d)
q, k, v = qkv.unbind(0)
# 使用 PyTorch 2.0 的 SDPA 加速
x = F.scaled_dot_product_attention(q, k, v, dropout_p=0.0)
x = x.transpose(1, 2).reshape(B, N, C)
return self.proj(x)

3.5 MLP 块#

class Mlp(nn.Module):
"""GELU MLP, hidden = 4 × in_dim。"""
def __init__(self, dim, mult=4):
super().__init__()
hidden = int(dim * mult)
self.fc1 = nn.Linear(dim, hidden)
self.act = nn.GELU(approximate="tanh")
self.fc2 = nn.Linear(hidden, dim)
def forward(self, x):
return self.fc2(self.act(self.fc1(x)))

4. AdaLN-Zero:条件注入的关键创新#

4.1 为什么 AdaLN-Zero 如此重要?#

扩散模型需要告诉网络”现在是去噪到哪一步了 tt”、“我要生成什么类别 yy”。

UNet 时代最常用的做法是 adaptive group normalization (AdaGN):把 t,yt, y 编码成一个向量,对每层归一化参数做调制。

DiT 论文研究了 4 种条件注入策略,发现最优方式是 AdaLN-Zero

4.2 四种条件注入策略#

论文的图如下(简化):

策略 1: In-context 策略 2: Cross-attention
───────────── ──────────────────
x → [t,y, x] → Transformer x → Transformer ──cross-attn── c
Q K V 从 x 来
条件直接拼成 token K, V 从 c 来
策略 3: Adaptive Layer-Norm 策略 4: adaLN-Zero ✅
────────────────── ─────────────────
x → LN → Modulated by c x → AdaLN(c) = γ·LN·x + α
→ Attention/MLP → Attention/MLP
(γ 和残差 α 都是 c 的函数)
初始化为零 → 残差为 0
(与残差连接相同效果先初始化)

4.3 AdaLN-Zero 的核心公式#

设条件嵌入 cc,对每个 Transformer 块计算 6 个调制参数:

γ1,γ2,β1,β2=MLP(c)R4D\gamma_1, \gamma_2, \beta_1, \beta_2 = \text{MLP}(c) \in \mathbb{R}^{4D}α1,α2=MLP(c)R2D\alpha_1, \alpha_2 = \text{MLP}(c) \in \mathbb{R}^{2D}

然后在每个残差块之前做调制:

h=a1+Attention(γ1LayerNorm(x)+β1)h = a_1 + \text{Attention}\left(\gamma_1 \cdot \text{LayerNorm}(x) + \beta_1\right)h=a2+MLP(γ2LayerNorm(h)+β2)h = a_2 + \text{MLP}\left(\gamma_2 \cdot \text{LayerNorm}(h) + \beta_2\right)

“Zero” 的含义:

  • 初始化 α\alpha 为零 → 整个块输出去残差效果 (identity)
  • 这样模型在一开始就是”什么都不做”的扩散初始化
  • 训练过程中,α\alpha 逐渐增长,模型逐渐学出能力

AdaLN-Zero 的三大好处

  1. 参数高效:6 维向量调制整个块,比 cross-attention 省得多
  2. 训练稳定:zero-init 让 loss 收敛更平顺
  3. 性能领先:在所有策略中 SOTA

4.4 完整 AdaLN DiT Block 代码#

class DiTBlock(nn.Module):
"""DiT 的核心 block: AdaLN-Zero 调制 + Self-Attention + MLP。"""
def __init__(self, hidden_size, num_heads, mlp_ratio=4.0):
super().__init__()
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
self.attn = Attention(hidden_size, num_heads=num_heads, qkv_bias=True)
self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
self.mlp = Mlp(hidden_size, mult=mlp_ratio)
# AdaLN 参数生成: 输入 c ∈ R^D, 输出 6 * D
self.adaLN_modulation = nn.Sequential(
nn.SiLU(),
nn.Linear(hidden_size, 6 * hidden_size, bias=True),
)
def forward(self, x, c):
"""
x: (B, S, D)
c: (B, D) 条件嵌入 (t + y)
"""
# 生成 6 个调制参数
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
self.adaLN_modulation(c).chunk(6, dim=1)
)
# 1. Attention 路径 (pre-LN + AdaLN + 门控残差)
x = x + gate_msa.unsqueeze(1) * self.attn(
modulate(self.norm1(x), shift_msa, scale_msa)
)
# 2. MLP 路径 (pre-LN + AdaLN + 门控残差)
x = x + gate_mlp.unsqueeze(1) * self.mlp(
modulate(self.norm2(x), shift_mlp, scale_mlp)
)
return x
def modulate(x, shift, scale):
"""AdaLN 调制: γ·x + β (但要求 γ, β 是 token-wise 的)。"""
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)

4.5 Zero-Initialization#

def _zero_init(module):
"""把 AdaLN 的输出初始化为 0 — 这一行很重要!"""
for p in module.parameters():
nn.init.zeros_(p)
# 在构造完 DiT 后调用一次
for block in model.blocks:
_zero_init(block.adaLN_modulation)
# final_layer 的 AdaLN 同样 zero init

5. 完整的 DiT 实现#

5.1 完整主类#

import math
import torch
import torch.nn as nn
import torch.nn.functional as F
class DiT(nn.Module):
"""Diffusion Transformer (class-conditional)."""
def __init__(
self,
img_size=32, patch_size=4, in_chans=3,
hidden_size=1152, depth=28, num_heads=16,
mlp_ratio=4.0, num_classes=1000,
):
super().__init__()
self.img_size = img_size
self.patch_size = patch_size
self.in_chans = in_chans
self.out_chans = in_chans # 预测噪声, 通道数相同
# 1) Patch 嵌入
self.patch_embed = PatchEmbed(img_size, patch_size, in_chans, hidden_size)
self.num_patches = self.patch_embed.num_patches
# 2) 位置编码 (可学习)
self.pos_embed = nn.Parameter(torch.zeros(1, self.num_patches, hidden_size))
# 3) 条件嵌入
self.t_embedder = TimestepEmbedder(hidden_size)
self.y_embedder = LabelEmbedder(num_classes, hidden_size)
# 4) Transformer blocks
self.blocks = nn.ModuleList([
DiTBlock(hidden_size, num_heads, mlp_ratio)
for _ in range(depth)
])
# 5) 最终 LayerNorm
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
self.adaLN_modulation_final = nn.Sequential(
nn.SiLU(),
nn.Linear(hidden_size, 2 * hidden_size, bias=True),
)
# 6) 线性解码器: 把 patch tokens → 噪声 patch
self.final_layer = nn.Linear(hidden_size, patch_size * patch_size * self.out_chans)
self.initialize_weights()
def initialize_weights(self):
# 1) 位置编码用 sin-cos 初始化
pos_embed = get_2d_sincos_pos_embed(
self.pos_embed.shape[-1], int(self.num_patches ** 0.5)
)
self.pos_embed.data.copy_(torch.from_numpy(pos_embed).float().unsqueeze(0))
# 2) 默认 Linears xavier
self.apply(self._init_weights)
# 3) AdaLN 调制层 ZERO init (关键的 AdaLN-Zero!)
for block in self.blocks:
nn.init.zeros_(block.adaLN_modulation[1].weight)
nn.init.zeros_(block.adaLN_modulation[1].bias)
nn.init.zeros_(self.adaLN_modulation_final[1].weight)
nn.init.zeros_(self.adaLN_modulation_final[1].bias)
# 4) final layer 归零
nn.init.zeros_(self.final_layer.weight)
nn.init.zeros_(self.final_layer.bias)
def _init_weights(self, m):
if isinstance(m, nn.Linear):
nn.init.xavier_uniform_(m.weight)
if m.bias is not None:
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.LayerNorm):
if m.weight is not None:
nn.init.constant_(m.weight, 1.0)
if m.bias is not None:
nn.init.constant_(m.bias, 0)
def unpatchify(self, x):
"""
x: (B, S, p*p*out_chans) → (B, C, H, W)
"""
B = x.shape[0]
p = self.patch_size
h = w = int(x.shape[1] ** 0.5)
x = x.reshape(B, h, w, p, p, self.out_chans)
x = torch.einsum("bhwpqc->bchpwq", x)
return x.reshape(B, self.out_chans, h * p, w * p)
def forward(self, x, t, y):
"""
x: (B, C, H, W) 噪声图像
t: (B,) 时间步
y: (B,) 类别标签
"""
# 1) 切块
x = self.patch_embed(x) # (B, S, D)
x = x + self.pos_embed # 加位置编码
# 2) 条件嵌入
c = self.t_embedder(t) + self.y_embedder(y)
# 3) 进入 N 个 DiT block
for block in self.blocks:
x = block(x, c)
# 4) 最终 AdaLN 调制
shift, scale = self.adaLN_modulation_final(c).chunk(2, dim=1)
x = modulate(self.norm_final(x), shift, scale)
# 5) 线性解码 → unpatchify → 噪声
x = self.final_layer(x) # (B, S, p*p*C)
x = self.unpatchify(x) # (B, C, H, W)
return x

5.2 DiT 模型规模表#

变体隐层 DD头数层数 NN参数量Gflops (256×256)
DiT-S38461233M5.9
DiT-B7681212130M23.0
DiT-L10241624458M80.8
DiT-XL11521628676M119.4
DiT-H14081632985M174.5

一个有意思的发现:DiT-S 大约等于原始 UNet 的计算量,但 FID 更低——Transformer 实际更高效。

6. 训练与采样#

6.1 训练目标#

DiT 沿用 DDPM/LDM 的标准训练目标(与 UNet 完全一致):

# === 标准 DDPM 训练循环 ===
def training_step(model, x0):
"""
x0: 真实图像 (B, C, H, W)
y: 类别标签 (B,)
"""
# 1. 随机时间步
t = torch.randint(0, num_train_timesteps, (x0.size(0),), device=x0.device)
# 2. 采样噪声
noise = torch.randn_like(x0)
# 3. 前向扩散 (重参数化)
xt = sqrt_alpha_bars[t, None, None, None] * x0 + \
sqrt_one_minus_alpha_bars[t, None, None, None] * noise
# 4. 预测噪声 (用 DiT 而非 UNet!)
pred_noise = model(xt, t, y)
# 5. MSE loss
loss = F.mse_loss(pred_noise, noise)
return loss

关键点:DiT 替代的只是 “pred_noise = …” 这一行,其他一切都不变。

6.2 采样 (Sampling)#

采样同样沿用 DDIM/DPM-Solver 流水线,无需修改(只要模型是噪声预测 ϵθ\epsilon_\theta):

@torch.no_grad()
def sample_ddim(model, shape, y, num_steps=50):
"""DDIM 采样 (50 步就足够)。"""
device = next(model.parameters()).device
# 准备时间步调度
ts = torch.linspace(num_train_timesteps - 1, 0, num_steps + 1).long().to(device)
x = torch.randn(*shape, device=device)
for i in range(num_steps):
t_cur = ts[i]
t_next = ts[i + 1]
# 预测噪声
pred_noise = model(x, t_cur.expand(x.size(0)), y)
# DDIM 更新公式
alpha_cur = alpha_bars[t_cur]
alpha_next = alpha_bars[t_next]
x0_pred = (x - (1 - alpha_cur).sqrt() * pred_noise) / alpha_cur.sqrt()
x = alpha_next.sqrt() * x0_pred + (1 - alpha_next).sqrt() * pred_noise
return x

6.3 Classifier-Free Guidance (CFG)#

DiT 与 LDMs 一样使用 CFG 来增强条件生成的视觉质量:

@torch.no_grad()
def sample_cfg(model, x, t, y, null_class, guidance_scale=4.0):
"""
CFG = unconditional + scale * (conditional - unconditional)
"""
eps_cond = model(x, t, y)
eps_uncond = model(x, t, null_class * torch.ones_like(y))
return eps_uncond + guidance_scale * (eps_cond - eps_uncond)

训练时需要有概率地把 yy 替换为 null 类别,让网络学会无条件生成:

def training_step_cfg(model, x0, y, uncond_prob=0.1):
t = torch.randint(0, num_train_timesteps, (x0.size(0),), device=x0.device)
# 10% 概率用 null 类别
drop_mask = torch.rand(y.shape).to(y.device) < uncond_prob
y_dropped = torch.where(drop_mask, null_class, y)
noise = torch.randn_like(x0)
xt = sqrt_alpha_bars[t] * x0 + sqrt_one_minus_alpha_bars[t] * noise
pred = model(xt, t, y_dropped)
return F.mse_loss(pred, noise)

7. 缩放律:DiT 最惊艳的结果#

7.1 Gflops vs FID 的近线性关系#

DiT 论文的一个核心贡献是发现:

FID(Fréchet Inception Distance,衡量生成质量的关键指标)随模型 Gflops 的增加,近似呈幂律下降——而且在极高 Gflops 时仍未饱和

模型GflopsFID ↓IS ↑
DiT-S (256)5.911.692.5
DiT-B (256)23.08.7124.0
DiT-L (256)80.85.7156.0
DiT-XL (256)119.44.5173.0
DiT-XL (512)525.02.3241.0
当时 SOTA (ADM-UNet)-4.0-

这条曲线意味着:只要继续把模型做大、做更长时间、喂更多数据,DiT 还能持续变好——没有”瓶颈”。

7.2 论文的缩放图#

FID
4 │ ● DiT-XL/512 (2.3)
│ ●
│ ●
8 │ ●
│ ●
│● DiT-S
└─────────────────────────────▶
5G 20G 80G 120G 500G
Gflops

7.3 与 UNet 的对比关键数据#

在相同的 Gflops 量级下:

相同 GflopsUNet (ADM)DiT差距
~80 GFID 4.0+FID 5.7DiT 以更低 FID 提升空间更大
缩放空间很快饱和持续下降DiT 拥有更好的缩放性

8. DiT 的扩展:从图像到视频、从文本到视觉#

8.1 Stable Diffusion 3 (MM-DiT)#

Stability AI 在 SD3 中采用了多模态 DiT (MM-DiT)

传统 DiT:
- 把文本 c 与图像 token 拼接 → AdaLN 调制
- 文本只通过 AdaLN 影响图像
MM-DiT (SD3):
- 文本 token 与图像 token 在每个 attention block 里
"joint attention" 一起做注意力
- 每个模态有独立的投影矩阵 (Q/K/V)
- 文本与图像"双向看": 文本看图像布局,图像看文本语义
┌─────────────────────────────────────────────┐
│ MM-DiT Block (SD3) │
│ │
│ img_tokens, text_tokens │
│ │ │
│ ▼ │
│ ┌──────────────────────────┐ │
│ │ LayerNorm │ │
│ │ Q = img_W_q @ img, ... │ 各自不同的 W │
│ │ K, V 拼接自 img + text │ │
│ │ Attn (Q, [K;V]) │ │
│ └──────────────────────────┘ │
│ │ │
│ ▼ │
│ MLP (各模态独立) │
└─────────────────────────────────────────────┘

8.2 Sora: 视频生成的基础#

2024 年 OpenAI 发布的 Sora 的核心架构正是 DiT 的视频版本

  • 把”图像 patch”换成”时空 patch” (spatial-temporal patches)
  • 一个 (H, W, T) 的视频被切成小块
  • DiT 在时空维度上做自注意力 → 全局理解每一帧
视频 patch 化:
原始视频 (T=16 帧, 256×256×3)
切成 16×16×2 时空小块 (2 帧, 16×16 像素)
得到 (T/2) × (H/16) × (W/16) 个 token
对于 16×256×256: 8 × 16 × 16 = 2048 个 token
与图像 DiT 一样的 Transformer 主干

Sora 的核心贡献之一:

只要 Transformer 主干够大、数据够多,DiT 可以学到世界模型 (world models)——视频里的物理因果关系。

8.3 与其他工作的关系谱系#

2022.12 DiT (Peebles & Xie)
2023 PixArt-α (轻量 DiT)
2023 SD3 / SDXL-Turbo (MM-DiT)
2024.2 Sora (视频 DiT)
2024.5 Hunyuan (腾讯视频 DiT)
2024-25 CogVideoX, Wan2.1 (开源视频 DiT)
│ Wan 2.1 16s 480P 单卡可跑
2025 Movie Gen, Veo 2 (电影级视频 DiT)

9. DiT 的工程实践#

9.1 加速技巧#

# 1) PyTorch 2.0 SDPA 加速注意力
torch.backends.cuda.sdp_kernel(enable_flash=True, enable_mem_efficient=True)
# 2) 混合精度
with torch.autocast(device_type='cuda', dtype=torch.bfloat16):
pred = model(x, t, y)
# 3) torch.compile
model = torch.compile(model, mode="reduce-overhead")

9.2 显存优化#

  • 梯度检查点 (Gradient Checkpointing): 把 DiT 的中间激活不存,只在反向时重算
  • FlashAttention: 长序列注意力降显存
  • 分桶并行 (Bucketed Parallel): DiT-S/B/L 用单 H100 即可训练;XL/H 需要张量并行
from torch.utils.checkpoint import checkpoint
class DiTBlock(nn.Module):
def forward(self, x, c):
# 用 checkpoint 让深度网络少占显存
return checkpoint(self._forward, x, c, use_reentrant=False)

9.3 Patch Size 的权衡#

情况小 patch (p=4)大 patch (p=8 / 16)
序列长度
计算量
局部精细度
推荐场景高分辨率低分辨率或大模型

9.4 类别 / 文本条件#

# 多类别 (含 null 类别用于 CFG)
self.y_embedder = nn.Embedding(num_classes + 1, hidden_size)
# 文本条件 (SD 时代)
self.text_embedder = ... # CLIP / T5 等文本编码器

10. DiT 与传统 UNet 的对比总结#

维度UNet (LDM/ADM)DiT
核心层2D Conv + ResNet blockSelf-Attn + MLP
归纳偏置强 (局部/平移不变)弱 (由数据学习)
缩放行为饱和较快严格单调下降
高分辨率多尺度金字塔天然适配需要 patch 选择 + patch 化
视频扩展需要 3D统一时空 patch 即可
与 LLM 融合模型架构不同架构一致 → 便于多模态
工程成熟度极其成熟仍在快速迭代
训练效率数据效率高数据量越大越有优势

DiT 的核心遗产:它确立了**“扩散 + Transformer”为生成模型的标准范式**——既拥抱 Transformer 的缩放性,又复用扩散模型的训练范式。

11. 局限与未来#

11.1 当前局限#

  • 数据饥渴:小数据上 DiT 比 UNet 表现差,需要 LAION/JFT 级的预训练
  • 高分辨率成本2562256^2 已是 1024 个 token;102421024^2 就是 16384 个 token,二次方的注意力很贵
  • 离散化训练:decoder 直接预测连续 patch 值,与 VAE decoder 配合时会有损失
  • CFG 仍是必需:多模态质量提升对 CFG scale 依赖大

11.2 未来方向#

加速方向:
├── Patch packing (更稀疏的 token)
├── Linear attention (Mamba-style) 替代 self-attention
├── 蒸馏到少步采样 (Consistency Model, SDXL-Turbo)
└── 潜在 patch 化 (更紧凑表示)
能力扩展:
├── 视频 DiT (Sora, Wan, CogVideoX)
├── 3D 视觉 (从图像扩散到 3D 资产生成)
├── 音频 DiT (MusicLDM, AudioLDM 2)
├── 多模态统一 DiT (同时文本/图/音/视频)
└── 具身控制 DiT (机器人动作生成)

11.3 一句话总结 DiT 的影响#

DiT 不只是替换了 UNet,它证明了生成模型的”图像 → 多模态”过渡应当建立在 Transformer 之上——这一判断现在被 Sora、Stable Diffusion 3、Gemini 视频等多模态生成系统反复印证。

12. 总结#

12.1 核心要点回顾#

维度关键要点
核心创新用 Transformer 替换 UNet 作为去噪网络
关键设计Patchify + AdaLN-Zero 调制
条件注入c=temb+yembc = t_{\text{emb}} + y_{\text{emb}},全块 AdaLN-Zero 调制
缩放行为FID vs Gflops 严格服从幂律下降
优势全局感受野、缩放性好、与 LLM 同架构
劣势数据饥渴、O(S2)O(S^2) 注意力高分辨率昂贵
影响Sora、SD3、Wan2.1 全部建立在 DiT 之上

12.2 推荐学习资源#

论文:
- DiT (2022): "Scalable Diffusion Models with Transformers" (Peebles & Xie)
- U-ViT (2022): "All are Worth Words"
- SD3 (2024): "Scaling Rectified Flow Transformers for High-Resolution Image Synthesis"
代码:
- facebookresearch/DiT (官方实现)
- HuggingFace diffusers (DiTPipeline)
- stabilityai/sd3 (MM-DiT 商业级)
扩展阅读:
- ViT (2020): Transformer 如何进入视觉
- DDPM / DDIM 原文: 扩散基础
- Sora 技术报告 (2024): 视频 DiT 巅峰

一句话总结:DiT 通过”patchify + AdaLN-Zero + 纯 Transformer 主干”这一简洁设计,把扩散模型推入 Transformer 时代,确立了图像/视频生成统一范式——也是 Sora 的技术源泉。

文章分享

如果这篇文章对你有帮助,欢迎分享给更多人!

深入理解 Diffusion Transformer (DiT):扩散模型的 Transformer 范式
https://aiattnstudio.link/posts/diffusion-transformer/
作者
Federico
发布于
2026-07-16
许可协议
CC BY-NC-SA 4.0
Profile Image of the Author

Federico

AI Research Lab

Hello, I'm Federico.

关于实验室 / About
公告

欢迎来到Federico的个人博客

分类
标签
站点统计
57文章
7分类
404标签