ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

深入理解 Diffusion Transformer (DiT):用 Transformer 替代 U-Net

深入理解 Diffusion Transformer (DiT):用 Transformer 替代 U-Net 1. 引言从 U-Net 到 Transformer1.1 U-Net 的统治在 DiT 出现之前扩散模型的骨干网络几乎被U-Net垄断。从 DDPM 到 Stable Diffusion去噪网络ϵ θ \epsilon_\thetaϵθ​无一例外地采用 U-Net 结构——一个带有编码器-解码器对称结构和**跳连skip connection**的卷积网络。U-Net 之所以流行是因为它的归纳偏置很适合图像去噪局部性卷积天然捕捉相邻像素的关系去噪本就是局部操作多尺度下采样-上采样结构能捕捉不同尺度的特征跳连把浅层细节直接传到深层适合保留高频细节。1.2 DiT 的疑问但 Peebles 和 Xie 在 2023 年 ICCV 论文“Scalable Diffusion Models with Transformers”中提出了一个尖锐的问题U-Net 的卷积归纳偏置真的是扩散模型所必需的吗此前ViTVision Transformer已经在分类、检测、分割等任务上证明Transformer 可以完全替代卷积且随规模增大表现更优。那么扩散模型的骨干网络是否也能换成 Transformer答案是肯定的——Diffusion TransformerDiT用一个纯 Transformer 替代了 U-Net并展示了更强的可扩展性scalability模型越大生成质量越高。这一思想后来成为Sora视频生成、Stable Diffusion 3、FLUX等前沿模型的架构基础。2. 背景在潜在空间里做扩散DiT 并非在像素空间直接扩散而是沿用了潜在扩散模型LDM的思路这能大幅降低计算量图像 (256×256×3) ──VAE 编码──► 潜向量 z (32×32×4) ──► DiT 去噪 ──► VAE 解码 ──► 生成图像用预训练的VAE把图像压缩 8 倍得到低维潜向量z ∈ R 32 × 32 × 4 z \in \mathbb{R}^{32 \times 32 \times 4}z∈R32×32×4DiT 只在这个潜空间里运行预测潜向量中的噪声ϵ θ ( z t , t , c ) \epsilon_\theta(z_t, t, c)ϵθ​(zt​,t,c)最后用 VAE 解码器把去噪后的潜向量还原成图像。训练目标就是熟悉的扩散损失L E t , z 0 , ϵ [ ∥ ϵ − ϵ θ ( z t , t , c ) ∥ 2 ] L \mathbb{E}_{t, z_0, \epsilon} \left[ \| \epsilon - \epsilon_\theta(z_t, t, c) \|^2 \right]LEt,z0​,ϵ​[∥ϵ−ϵθ​(zt​,t,c)∥2]其中c cc是条件如类别标签。3. DiT 的整体架构Patchify3.1 把潜向量切成 PatchTransformer 只能处理序列而潜向量是二维的。DiT 借鉴 ViT 的做法把潜向量切成一个个不重叠的patch潜向量 z (32×32×4) │ Patchify (patch size p2) ▼ ┌───┬───┬───┬───┐ 每个 patch: 2×2×4 16 维 │ 1 │ 2 │ 3 │ 4 │ 共 T (32/2)² 256 个 token ├───┼───┼───┼───┤ │ 5 │ 6 │ 7 │ 8 │ ├───┼───┼───┼───┤ │ . │ . │ . │ . │ ├───┼───┼───┼───┤ │...│...│...│...│ └───┴───┴───┴───┘ │ 每个 patch 拉平后线性投影到维度 d ▼ token 序列: [z₁, z₂, ..., z_T], zᵢ ∈ ℝᵈ具体地若潜向量尺寸为I × I × C I \times I \times CI×I×C、patch 尺寸为p pp则token 数量T ( I / p ) 2 T (I / p)^2T(I/p)2每个 token 的原始维度p 2 ⋅ C p^2 \cdot Cp2⋅C通过一个线性层投影到 Transformer 的隐藏维度d dd。DiT 论文中256×256 图像经 VAE 压缩为32 × 32 × 4 32\times32\times432×32×4的潜向量patch sizep 2 p2p2得到 256 个 token隐藏维度d 1152 d1152d1152。3.2 位置编码与 ViT 一样DiT 使用正弦位置编码二维注入位置信息。随后 token 序列经过N NN个 DiT Block 处理最后用一个线性层和 reshape 还原成去噪后的潜向量。潜向量 z_t ──Patchify──► token 序列 位置编码 ──► [ DiT Block × N ] ──► Linear Unpatchify ──► 预测噪声 ▲ 条件 c (类别 时间步) ┘4. 条件注入的四种方式如何把时间步t tt和类别标签y yy统称条件c cc注入 Transformer是 DiT 重点消融的设计。论文对比了四种方案方案做法特点In-context把条件 token 拼到输入序列末尾随序列一起进 Transformer最简单但占用序列长度Cross-attention把条件拼成序列作为交叉注意力的 Key/Value增加计算需额外注意力层adaLN条件回归出 LayerNorm 的 scale/shift逐层调制轻量、有效adaLN-Zero在 adaLN 基础上把残差门控初始化为 0最优实验结论清晰adaLN-Zero 效果最好且比 cross-attention 更高效不需要额外的注意力计算。5. adaLN-Zero让每个块从恒等映射起步5.1 自适应层归一化adaLN标准 LayerNorm 是无参数的固定归一化。adaLN让条件c cc通过一个小 MLP回归出每个维度的缩放α \alphaα和偏移β \betaβadaLN ( h , α , β ) LayerNorm ( h ) ⊙ ( 1 α ) β \text{adaLN}(h, \alpha, \beta) \text{LayerNorm}(h) \odot (1 \alpha) \betaadaLN(h,α,β)LayerNorm(h)⊙(1α)β其中α , β ∈ R d \alpha, \beta \in \mathbb{R}^{d}α,β∈Rd是逐维度的调制参数1 α 1 \alpha1α保证初始时α 0 \alpha0α0退化为标准 LayerNorm。5.2 adaLN-Zero 的关键改进adaLN-Zero 在此基础上加了两个门控参数γ \gammaγ分别作用于注意力分支和 MLP 分支的残差输出h ← h γ 1 ⊙ Attention ( adaLN ( h , α 1 , β 1 ) ) h \leftarrow h \gamma_1 \odot \text{Attention}\big(\text{adaLN}(h, \alpha_1, \beta_1)\big)h←hγ1​⊙Attention(adaLN(h,α1​,β1​))h ← h γ 2 ⊙ MLP ( adaLN ( h , α 2 , β 2 ) ) h \leftarrow h \gamma_2 \odot \text{MLP}\big(\text{adaLN}(h, \alpha_2, \beta_2)\big)h←hγ2​⊙MLP(adaLN(h,α2​,β2​))最关键的是初始化技巧把回归 MLP 的最后一层权重和偏置初始化为 0使得初始时α β γ 0 \alpha\beta\gamma0αβγ0。于是h ← h 0 h h \leftarrow h 0 hh←h0h直觉初始时每个 DiT Block 都是恒等映射整个 DiT 一开始就是一个直通网络条件的影响从零逐步长出来。这种残差门控类似 ReZero让深层 Transformer 的训练极其稳定也加速了收敛。5.3 一个 DiT Block 的完整流程条件 c (类别 时间步) │ ┌────────▼────────┐ │ MLP (SiLU→Linear)│ 输出 6 组参数 │ α1 β1 γ1 α2 β2 γ2│ └────────┬────────┘ │ ┌──────────────┴───────────────┐ │ │ ┌────▼────┐ ┌────▼────┐ │ LayerNorm│ │ LayerNorm│ │ ×(1α1) │ │ ×(1α2) │ │ β1 │ │ β2 │ └────┬────┘ └────┬────┘ │ │ ┌────▼────┐ ┌────▼────┐ │ Multi-Head│ │ MLP │ │ Attention│ │ │ └────┬────┘ └────┬────┘ │ ×γ1 │ ×γ2 └──────────► h ◄─────────────┘ (残差相加h γ1·attn γ2·mlp)6. 缩放性DiT 的真正价值6.1 计算量随规模增长DiT 最核心的实验发现是缩放定律随着模型计算量Gflops增加生成质量FID持续提升。模型Patch size隐藏维度层数参数GflopsDiT-S23841233M6.85DiT-B276812130M23.01DiT-L2102424458M80.71DiT-XL2115228675M118.646.2 关键结论模型越大FID 越低DiT-XL/2 在 ImageNet 256×256 上取得 FID 2.27使用无分类器引导超越同期 GAN 与 U-Net 扩散模型Transformer 无卷积归纳偏置也能工作证明 U-Net 的卷积结构并非必需注意力机制足以胜任去噪可扩展性优于 U-Net在相同计算量下DiT 比 U-Net 骨干扩展得更稳、更高效。深层意义DiT 把扩散模型的骨干网络与 ViT 对齐从而让扩散模型也能搭上 Transformer 缩放定律的快车——这是后续 Sora 等视频生成模型能够大规模扩展的关键前提。7. PyTorch 实现下面给出 DiT 核心组件的精简实现Patchify、adaLN-Zero 调制、DiT Block。7.1 调制与 Patchifyimporttorchimporttorch.nnasnnimportmathdefmodulate(x,shift,scale):adaLNLayerNorm 后做逐维度缩放与偏移returnx*(1scale.unsqueeze(1))shift.unsqueeze(1)classPatchify(nn.Module):把潜向量 (B, C, H, W) 切成 patch 序列 (B, T, d)def__init__(self,patch_size,in_ch,hidden_size):super().__init__()self.patch_sizepatch_size self.projnn.Linear(patch_size*patch_size*in_ch,hidden_size)defforward(self,x):B,C,H,Wx.shape pself.patch_size# (B, C, H, W) - (B, T, p*p*C)xx.reshape(B,C,H//p,p,W//p,p)xx.permute(0,2,4,1,3,5).reshape(B,-1,p*p*C)returnself.proj(x)7.2 位置编码defget_2d_sincos_pos_embed(hidden_size,grid_size):二维正弦位置编码返回 (grid_size², hidden_size)grid_htorch.arange(grid_size,dtypetorch.float32)grid_wtorch.arange(grid_size,dtypetorch.float32)gridtorch.stack(torch.meshgrid(grid_h,grid_w,indexingij),dim0)# (2, H, W)gridgrid.reshape(2,1,grid_size*grid_size)# (2, 1, T)emb_hget_1d_sincos(grid[0],hidden_size//2)# (T, d/2)emb_wget_1d_sincos(grid[1],hidden_size//2)returntorch.cat([emb_h,emb_w],dim1)# (T, d)defget_1d_sincos(pos,dim):freqstorch.exp(-math.log(10000)*torch.arange(0,dim,2)/dim)argspos.reshape(-1,1)*freqs# (T, dim/2)returntorch.cat([torch.sin(args),torch.cos(args)],dim1)7.3 adaLN-Zero 的 DiT BlockclassDiTBlock(nn.Module):def__init__(self,hidden_size,num_heads,mlp_ratio4.0):super().__init__()self.norm1nn.LayerNorm(hidden_size,elementwise_affineFalse)self.attnnn.MultiheadAttention(hidden_size,num_heads,batch_firstTrue)self.norm2nn.LayerNorm(hidden_size,elementwise_affineFalse)self.mlpnn.Sequential(nn.Linear(hidden_size,int(hidden_size*mlp_ratio)),nn.GELU(approximatetanh),nn.Linear(int(hidden_size*mlp_ratio),hidden_size),)# 由条件 c 回归 6 组参数: shift_msa, scale_msa, gate_msa,# shift_mlp, scale_mlp, gate_mlpself.adaLNnn.Sequential(nn.SiLU(),nn.Linear(hidden_size,6*hidden_size))# 关键最后一层零初始化使 block 初始为恒等映射nn.init.constant_(self.adaLN[1].weight,0)nn.init.constant_(self.adaLN[1].bias,0)defforward(self,x,c):shift_msa,scale_msa,gate_msa,shift_mlp,scale_mlp,gate_mlp\ self.adaLN(c).chunk(6,dim1)# 注意力分支hmodulate(self.norm1(x),shift_msa,scale_msa)attn_out,_self.attn(h,h,h)xxgate_msa.unsqueeze(1)*attn_out# MLP 分支hmodulate(self.norm2(x),shift_mlp,scale_mlp)xxgate_mlp.unsqueeze(1)*self.mlp(h)returnx7.4 时间步与类别条件编码classTimestepEmbedder(nn.Module):把时间步 t 编码为高维向量正弦 MLPdef__init__(self,hidden_size):super().__init__()self.mlpnn.Sequential(nn.Linear(hidden_size,hidden_size*4),nn.SiLU(),nn.Linear(hidden_size*4,hidden_size),)defforward(self,t):halft.shape[-1]freqstorch.exp(-math.log(10000)*torch.arange(half,devicet.device)/half)argst.unsqueeze(-1)*freqs embtorch.cat([torch.sin(args),torch.cos(args)],dim-1)returnself.mlp(emb)完整复现可参考官方 facebookresearch/DiT。上述代码体现了 DiT 的两个灵魂设计Patchify 把潜向量变成 token 序列adaLN-Zero 用零初始化让每个块从恒等映射起步。8. 影响与应用DiT 的影响远超一篇论文它奠定了新一代扩散模型的架构方向Sora2024OpenAI 的视频生成模型将 DiT 扩展到时空潜空间用 Transformer 统一处理视频 patchStable Diffusion 3 / FLUX新一代文生图模型采用 DiT 架构 Flow Matching / Rectified Flow 训练统一多模态骨干DiT 证明 Transformer 能统一图像、视频的生成为一个 Transformer 处理一切的愿景铺路架构研究催生了关于扩散模型缩放定律、patch 尺寸选择、条件注入方式的大量后续研究。9. 总结DiT 用一次骨干网络的替换实验回答了一个深刻的问题扩散模型到底需要什么样的网络结构Patchify把潜空间图像切成 token让 Transformer 进入扩散模型adaLN-Zero用条件回归调制参数 零初始化门控让深层 Transformer 训练稳定、收敛更快缩放性模型越大生成质量越高DiT 因此成为大规模生成模型的首选骨干深远影响直接催生了 Sora、SD3、FLUX 等前沿工作。DiT 的启示在于当你不确定某个结构是否必需时去掉它、换成更通用的方案然后交给规模去验证。正是这种去归纳偏置的勇气让扩散模型从 U-Net 时代跨入了 Transformer 时代。参考文献Peebles, W., Xie, S.Scalable Diffusion Models with Transformers.ICCV 2023.Dosovitskiy, A., et al.An Image is Worth 16x16 Words: Transformers for Image Recognition at Scale.ICLR 2021.Ho, J., Jain, A., Abbeel, P.Denoising Diffusion Probabilistic Models.NeurIPS 2020.Rombach, R., et al.High-Resolution Image Synthesis with Latent Diffusion Models.CVPR 2022.Brooks, T., et al.Video Generation Models as World Simulators.OpenAI 2024.Esser, P., et al.Scaling Rectified Flow Transformers for High-Resolution Image Synthesis.ICML 2024.
返回列表