深度理解-DiT
·
Diffusion 原理:
https://www.zhihu.com/tardis/zm/art/599887666?source_id=1005
先从一个宏观的角度去总结Diffusion的原理:
- 前向(扩散)过程:通过给原始图片添加随机噪声和时间,来训练噪声预测模型(对应模型架构中的Unet)
- 反向(推理)过程:将带有噪声的图片还原为清晰的图片,或者由原始的随机噪声,通过引导比如文字引导其生成对应的图像
1 前向扩散过程
输入:一张清晰的原始图片 x0
将时间定义为从 0 到 T 的离散步长(例如 T=1000)。每一步 t,我们都向上一步的图像中混入少量的高斯噪声
xt=1−βtxt−1+βtϵx_t = \sqrt{1-\beta_t}x_{t-1} + \sqrt{\beta_t}\epsilonxt=1−βtxt−1+βtϵ
其中 ϵ\epsilonϵ 是本次添加的随机噪声。
βt\beta_tβt 是一个预先设定好的“噪声调度器”(Noise Scheduler),控制每一步加噪的比例。随着 t 增加,原始信号 x0 的占比越来越小,噪声 ϵ\epsilonϵ的占比越来越大
最终状态:当 t=T时,图像 xT基本上就变成了纯粹的随机高斯噪声(标准正态分布),原始图像信息几乎完全丢失。
2 训练过程 U-Net U-Net的结构这里不过多展开
输入:
- 带噪图像 xtx_txt(维度例如 [batch, 3, 64, 64])。
- 当前时间步 ttt(告诉模型现在的噪声强度大概是多少)。
- (可选) 文本提示词的 Embedding(用于引导生成)。
输出: - 预测的噪声 ϵ^(维度与输入图像 xtx_txt 相同,即 [batch, 3, 64, 64])。
损失函数(Loss Function):计算真实噪声 ϵ和模型预测的噪声 ϵ^ 之间的差异
3 反向推理过程:
初始状态:从纯粹的随机噪声开始,假设它就是 T 时刻的图像 xT
迭代去噪(Iterative Denoising):
- 这是一个循环过程,从 t=T逐步走向 t=0
- 在每一步 t:
- 我们将当前的噪点图 xt、时间 t(以及文本条件)喂给训练好的 U-Net。
- U-Net 预测出当前图像里包含的噪声 ϵ^。

- 注意:这里不是一次性减去所有噪声,而是根据噪声调度器,只减去当前步骤对应的那一部分。
最终结果:经过 T 次迭代(或者使用加速采样器如 DDIM、DPM-Solver 只需几十次迭代),当时间到达 t=0 时,我们就得到了一张清晰的、符合数据分布的图像。
下面放一个详细的模型结构:展示了一个model如何根据输入文字生成图片
更详细的推理可以去看https://www.zhihu.com/tardis/zm/art/599887666?source_id=1005
DiT原理:
DiT 论文表明 U-Net 架构设计对 Diffusion Models 的性能并不重要,并且它们可以很容易地替换为 Transformers,这能带来更高的可扩展性,鲁棒性
结构解读:
- Patchify:和ViT类似,划分为若干patchs,将图像转换为序列
- DiT Block:除了噪声图像输入之外,有时会处理额外的条件信息,比如噪声时间步长,类标签,自然语言或者是别的隐变量(latent token)不必拘泥于上述的几个东西
这里展示了Block的三种形式:
a. In-Context Conditioning: 只需要将时间步长 , 类标签 作为2个额外的 token 拼接到输入序列中,类似于 ViT 的 [CLS] token
b. Cross Attention Block:给 Transformer Block 添加一个 Cross-Attention 块,时间步 t 和类别标签 c(或文本提示)被编码后,作为 Key ( K ) 和 Value ( V ) 输入到 Cross-Attention 模块中,但因为多了整个一层 Attention,计算量大
c. adaLN-Zero Block:除了回归计算缩放和移位参数 γ\gammaγ 和 β\betaβ 之外,还引入门控 α\alphaα 计算公式:y=Attention(Norm(x)⋅(1+γ)+β)⋅α+x,其数学本质在于动态改变特征分布
(Zero-Init α;强制初始化 α=0,解决最初随机初始化的缩放和移位参数 γ\gammaγ 和 β\betaβ 不稳定,无法收敛的问题) - Transformer Decoder:在最后一个 DiT Block 之后,需要将 image tokens 的序列解码为输出 噪声 以及 协方差的预测结果

手撕DiTBlock:
import torch
from torch import nn
from timm.models.vision_transformer import Attention, Mlp
def modulate(x, shift, scale):
return x*(1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
class DiTBlock(nn.Module):
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)
mlp_hidden_dim = int(hidden_size*mlp_ratio)
approx_gelu = lambda: nn.GELU(approximate="tanh")
self.mlp = Mlp(in_features=hidden_size, hidden_features=mlp_hidden_dim, out_features=hidden_size, act_layer=approx_gelu)
self.adaLN_modulation = nn.Sequential(
nn.SiLU(),
nn.Linear(hidden_size, 6*hidden_size, bias=True)
)
def forward(self, x, c):
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=1)
x = x + self.attn(modulate(self.norm1(x), shift_msa, scale_msa))*gate_msa.unsequeeze(1)
x = x + self.mlp(modulate(self.norm2(x), shift_mlp, scale_mlp))*gate_mlp.unsequeeze(1)
return x
AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。
更多推荐



所有评论(0)