深入理解 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 对比
| 维度 | UNet | DiT |
|---|---|---|
| 核心算子 | 2D 卷积 (ResNet block) | Self-Attention + MLP |
| 归纳偏置 | 强 (局部、平移不变) | 弱 (全局,由数据学习) |
| 信息传递 | 跳层连接 + 多尺度 | 全局自注意力 |
| 时间/类别条件 | 拼接或 AdaGN | AdaLN-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 的第一步是把 的图像转成 token 序列(与 ViT 类似)。设 patch 大小为 ,则序列长度:
| 输入尺寸 | Patch | 序列长度 |
|---|---|---|
| 32×32 | 4 | 64 |
| 64×64 | 4 | 256 |
| 256×256 | 8 | 1024 |
| 256×256 | 16 | 256 |
import torchimport 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 x3.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 相加,得到最终条件向量 :
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 如此重要?
扩散模型需要告诉网络”现在是去噪到哪一步了 ”、“我要生成什么类别 ”。
UNet 时代最常用的做法是 adaptive group normalization (AdaGN):把 编码成一个向量,对每层归一化参数做调制。
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 的核心公式
设条件嵌入 ,对每个 Transformer 块计算 6 个调制参数:
然后在每个残差块之前做调制:
“Zero” 的含义:
- 初始化 为零 → 整个块输出去残差效果 (identity)
- 这样模型在一开始就是”什么都不做”的扩散初始化
- 训练过程中, 逐渐增长,模型逐渐学出能力
AdaLN-Zero 的三大好处:
- 参数高效:6 维向量调制整个块,比 cross-attention 省得多
- 训练稳定:zero-init 让 loss 收敛更平顺
- 性能领先:在所有策略中 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 init5. 完整的 DiT 实现
5.1 完整主类
import mathimport torchimport torch.nn as nnimport 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 x5.2 DiT 模型规模表
| 变体 | 隐层 | 头数 | 层数 | 参数量 | Gflops (256×256) |
|---|---|---|---|---|---|
| DiT-S | 384 | 6 | 12 | 33M | 5.9 |
| DiT-B | 768 | 12 | 12 | 130M | 23.0 |
| DiT-L | 1024 | 16 | 24 | 458M | 80.8 |
| DiT-XL | 1152 | 16 | 28 | 676M | 119.4 |
| DiT-H | 1408 | 16 | 32 | 985M | 174.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 流水线,无需修改(只要模型是噪声预测 ):
@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 x6.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)训练时需要有概率地把 替换为 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 时仍未饱和。
| 模型 | Gflops | FID ↓ | IS ↑ |
|---|---|---|---|
| DiT-S (256) | 5.9 | 11.6 | 92.5 |
| DiT-B (256) | 23.0 | 8.7 | 124.0 |
| DiT-L (256) | 80.8 | 5.7 | 156.0 |
| DiT-XL (256) | 119.4 | 4.5 | 173.0 |
| DiT-XL (512) | 525.0 | 2.3 | 241.0 |
| 当时 SOTA (ADM-UNet) | - | 4.0 | - |
这条曲线意味着:只要继续把模型做大、做更长时间、喂更多数据,DiT 还能持续变好——没有”瓶颈”。
7.2 论文的缩放图
FID │4 │ ● DiT-XL/512 (2.3) │ ● │ ●8 │ ● │ ● │● DiT-S └─────────────────────────────▶ 5G 20G 80G 120G 500G Gflops7.3 与 UNet 的对比关键数据
在相同的 Gflops 量级下:
| 相同 Gflops | UNet (ADM) | DiT | 差距 |
|---|---|---|---|
| ~80 G | FID 4.0+ | FID 5.7 | DiT 以更低 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.compilemodel = 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 block | Self-Attn + MLP |
| 归纳偏置 | 强 (局部/平移不变) | 弱 (由数据学习) |
| 缩放行为 | 饱和较快 | 严格单调下降 |
| 高分辨率 | 多尺度金字塔天然适配 | 需要 patch 选择 + patch 化 |
| 视频扩展 | 需要 3D | 统一时空 patch 即可 |
| 与 LLM 融合 | 模型架构不同 | 架构一致 → 便于多模态 |
| 工程成熟度 | 极其成熟 | 仍在快速迭代 |
| 训练效率 | 数据效率高 | 数据量越大越有优势 |
DiT 的核心遗产:它确立了**“扩散 + Transformer”为生成模型的标准范式**——既拥抱 Transformer 的缩放性,又复用扩散模型的训练范式。
11. 局限与未来
11.1 当前局限
- 数据饥渴:小数据上 DiT 比 UNet 表现差,需要 LAION/JFT 级的预训练
- 高分辨率成本: 已是 1024 个 token; 就是 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 调制 |
| 条件注入 | ,全块 AdaLN-Zero 调制 |
| 缩放行为 | FID vs Gflops 严格服从幂律下降 |
| 优势 | 全局感受野、缩放性好、与 LLM 同架构 |
| 劣势 | 数据饥渴、 注意力高分辨率昂贵 |
| 影响 | 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 的技术源泉。
文章分享
如果这篇文章对你有帮助,欢迎分享给更多人!

