扩散模型MNIST实战

扩散模型学习笔记:从 DDPM 到 FLUX(MNIST 实战)

本文档整合扩散模型发展的三个过程,分别实现MNIST数据集的生成,构成一条完整的扩散模型学习路径:

  • 第一部分:基础扩散模型 —— DDPM 原理、MNIST 像素空间生成、DDIM 加速采样
  • 第二部分:Mini Stable Diffusion —— 潜空间扩散、条件注入、Classifier-Free Guidance
  • 第三部分:Mini FLUX —— Flow Matching、Euler ODE 采样、双流 MMDiT

三部分共用同一套 MNIST 实验平台,每一部分只在前一部分的基础上做增量改动,
可以清晰看到从 DDPM 到 Stable Diffusion 再到 FLUX 的演进逻辑。
所有代码均为从零手写、可直接运行,不依赖任何已有项目实现。环境要求:torch、torchvision。

目录


第一部分:DDPM 与 DDIM

1. Diffusion 模型的核心思想

扩散模型是一类生成模型,灵感来自非平衡热力学。它的核心思想可以用一句话概括:

学习"如何加噪"的逆过程——“如何一步步去噪”,从而把纯噪声还原成真实数据。

整个框架由两个过程组成:

  1. 前向过程(Forward / Diffusion Process):固定的马尔可夫链,逐步向数据中加入高斯噪声,直到数据变成纯噪声。这个过程不需要学习。
  2. 反向过程(Reverse Process):学习的马尔可夫链,从纯噪声出发,逐步去噪,最终生成新的数据样本。训练的目标就是学会这个去噪过程。
1
2
前向加噪:  x_0 → x_1 → x_2 → ... → x_T (≈ 纯高斯噪声)
反向去噪: x_T → x_{T-1} → ... → x_1 → x_0 (≈ 生成的样本)

2. 前向过程

2.1 单步加噪

给定真实样本 x0∼q(x0)x_0 \sim q(x_0),前向过程在每个时间步加入少量高斯噪声:

q(xt∣xt−1)=N(xt; 1−βt xt−1, βtI)q(x_t \mid x_{t-1}) = \mathcal{N}\left(x_t;\ \sqrt{1-\beta_t}\,x_{t-1},\ \beta_t \mathbf{I}\right)

其中 βt∈(0,1)\beta_t \in (0,1) 是预设的噪声强度(noise schedule),通常满足
β1<β2<⋯<βT\beta_1 < \beta_2 < \dots < \beta_T,即越往后加的噪声越多。

等价地用重参数化写法:

xt=1−βt xt−1+βt ϵt−1,ϵt−1∼N(0,I)x_t = \sqrt{1-\beta_t}\,x_{t-1} + \sqrt{\beta_t}\,\epsilon_{t-1},\qquad \epsilon_{t-1}\sim\mathcal{N}(0,\mathbf{I})

2.2 任意时间步的闭式表达

这是 DDPM 最重要的性质之一:不需要逐步迭代,可以一步到位从 x0x_0 直接采样任意 tt 时刻的 xtx_t。

定义 αt=1−βt\alpha_t = 1-\beta_t,αˉt=∏s=1tαs\bar{\alpha}_t = \prod_{s=1}^{t}\alpha_s。由于两个独立高斯噪声叠加仍然是高斯噪声(方差相加),递推可得:

q(xt∣x0)=N(xt; αˉt x0, (1−αˉt)I)q(x_t \mid x_0) = \mathcal{N}\left(x_t;\ \sqrt{\bar{\alpha}_t}\,x_0,\ (1-\bar{\alpha}_t)\mathbf{I}\right)

即:

xt=αˉt x0+1−αˉt ϵ,ϵ∼N(0,I)\boxed{x_t = \sqrt{\bar{\alpha}_t}\,x_0 + \sqrt{1-\bar{\alpha}_t}\,\epsilon,\qquad \epsilon\sim\mathcal{N}(0,\mathbf{I})}

当 TT 足够大时 αˉT→0\bar{\alpha}_T \to 0,xTx_T 近似为标准高斯分布——这正是采样时的起点。

这个公式同时给出了训练数据的构造方式:随机取一个 tt,加噪得到 xtx_t,就得到了一个训练样本。

2.3 噪声调度(Noise Schedule)

βt\beta_t 的取法有很多,最经典的是线性调度:

βt=βmin⁡+tT(βmax⁡−βmin⁡)\beta_t = \beta_{\min} + \frac{t}{T}\left(\beta_{\max}-\beta_{\min}\right)

常用取值:T=1000T=1000,βmin⁡=10−4\beta_{\min}=10^{-4},βmax⁡=0.02\beta_{\max}=0.02。

3. 反向过程

3.1 理想情况

如果能直接知道 q(xt−1∣xt)q(x_{t-1}\mid x_t),就可以从 xT∼N(0,I)x_T\sim\mathcal{N}(0,\mathbf{I}) 逐步采样回 x0x_0。可以证明,在已知 x0x_0 的条件下,反向转移核同样是高斯分布:

q(xt−1∣xt,x0)=N(xt−1; μ~t(xt,x0), β~tI)q(x_{t-1}\mid x_t, x_0) = \mathcal{N}\left(x_{t-1};\ \tilde{\mu}_t(x_t,x_0),\ \tilde{\beta}_t\mathbf{I}\right)

其中

β~t=1−αˉt−11−αˉt βt\tilde{\beta}_t = \frac{1-\bar{\alpha}_{t-1}}{1-\bar{\alpha}_t}\,\beta_t

μ~t(xt,x0)=αˉt−1 βt1−αˉt x0+αt (1−αˉt−1)1−αˉt xt\tilde{\mu}_t(x_t,x_0) = \frac{\sqrt{\bar{\alpha}_{t-1}}\,\beta_t}{1-\bar{\alpha}_t}\,x_0 + \frac{\sqrt{\alpha_t}\,(1-\bar{\alpha}_{t-1})}{1-\bar{\alpha}_t}\,x_t

问题在于:采样时我们没有 x0x_0(否则就不用生成了)。因此需要用一个神经网络 pθp_\theta 来近似这个转移分布。

3.2 用神经网络近似

DDPM 将反向过程建模为:

pθ(xt−1∣xt)=N(xt−1; μθ(xt,t), σt2I)p_\theta(x_{t-1}\mid x_t) = \mathcal{N}\left(x_{t-1};\ \mu_\theta(x_t,t),\ \sigma_t^2\mathbf{I}\right)

方差 σt2\sigma_t^2 直接取固定值 β~t\tilde{\beta}_t,网络只需要学习均值 μθ\mu_\theta。

3.3 从"预测均值"到"预测噪声"

直接预测均值可行但不稳定。DDPM 论文的关键技巧是把均值的参数化方式改一下。由前向公式反解:

x0=xt−1−αˉt ϵαˉtx_0 = \frac{x_t - \sqrt{1-\bar{\alpha}_t}\,\epsilon}{\sqrt{\bar{\alpha}_t}}

代入 μ~t\tilde{\mu}_t 的表达式化简可得:

μ~t=1αt(xt−βt1−αˉt ϵ)\tilde{\mu}_t = \frac{1}{\sqrt{\alpha_t}}\left(x_t - \frac{\beta_t}{\sqrt{1-\bar{\alpha}_t}}\,\epsilon\right)

也就是说,均值完全由噪声 ϵ\epsilon 决定。于是我们让网络 ϵθ(xt,t)\epsilon_\theta(x_t,t) 去预测加入的噪声,均值即为:

μθ(xt,t)=1αt(xt−βt1−αˉt ϵθ(xt,t))\mu_\theta(x_t,t) = \frac{1}{\sqrt{\alpha_t}}\left(x_t - \frac{\beta_t}{\sqrt{1-\bar{\alpha}_t}}\,\epsilon_\theta(x_t,t)\right)

3.4 训练目标

生成模型的标准目标是最大化似然 pθ(x0)p_\theta(x_0)。对扩散模型,可以通过变分下界(ELBO)推导,化简后的训练损失惊人地简单——就是让网络预测的噪声与真实加入的噪声之间的均方误差:

Lsimple=Ex0, t, ϵ∥ϵ−ϵθ(αˉt x0+1−αˉt ϵ, t)∥2\boxed{\mathcal{L}_{\text{simple}} = \mathbb{E}_{x_0,\,t,\,\epsilon}\left\|\epsilon - \epsilon_\theta\left(\sqrt{\bar{\alpha}_t}\,x_0 + \sqrt{1-\bar{\alpha}_t}\,\epsilon,\ t\right)\right\|^2}

训练算法(每次迭代):

  1. 从数据集中采样 x0x_0
  2. 随机采样时间步 t∼Uniform{1,…,T}t \sim \text{Uniform}\{1,\dots,T\}
  3. 采样噪声 ϵ∼N(0,I)\epsilon \sim \mathcal{N}(0,\mathbf{I})
  4. 计算 xt=αˉt x0+1−αˉt ϵx_t = \sqrt{\bar{\alpha}_t}\,x_0 + \sqrt{1-\bar{\alpha}_t}\,\epsilon
  5. 以 ∥ϵ−ϵθ(xt,t)∥2\|\epsilon - \epsilon_\theta(x_t, t)\|^2 为损失做梯度下降

采样算法(生成新图像):

  1. 采样 xT∼N(0,I)x_T \sim \mathcal{N}(0,\mathbf{I})
  2. 对 t=T,T−1,…,1t = T, T-1, \dots, 1:

    xt−1=1αt(xt−βt1−αˉt ϵθ(xt,t))+σt z,z∼N(0,I)x_{t-1} = \frac{1}{\sqrt{\alpha_t}}\left(x_t - \frac{\beta_t}{\sqrt{1-\bar{\alpha}_t}}\,\epsilon_\theta(x_t,t)\right) + \sigma_t\,z,\qquad z\sim\mathcal{N}(0,\mathbf{I})

    (t=1t=1 时不加噪声项)
  3. 返回 x0x_0

4. 网络结构:带时间条件的 U-Net

ϵθ(xt,t)\epsilon_\theta(x_t, t) 的输入和输出都是图像形状,同时还需要注入时间步 tt 的信息。标准选择是 U-Net:

  • 下采样路径:卷积逐步提取特征、降低分辨率
  • 上采样路径:转置卷积/插值逐步恢复分辨率
  • 跳跃连接(skip connection):把下采样各层的特征拼到对应的上采样层,保留细节信息
  • 时间嵌入:将 tt 用正弦位置编码(Sinusoidal Embedding,与 Transformer 相同)映射为向量,再通过 MLP 投影后加到每个卷积块的特征上,让网络知道"当前处于去噪的第几步"

正弦时间嵌入的公式(dd 为嵌入维度):

emb(t)2i=sin⁡(t100002i/d),emb(t)2i+1=cos⁡(t100002i/d)\text{emb}(t)_{2i} = \sin\left(\frac{t}{10000^{2i/d}}\right),\qquad \text{emb}(t)_{2i+1} = \cos\left(\frac{t}{10000^{2i/d}}\right)

5. DDPM 的 MNIST 完整实现

下面是一个完整的、可直接运行的实现,包含训练与采样。环境要求:torch、torchvision。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
# ddpm_mnist.py
# 从零实现的 DDPM,在 MNIST 上训练并生成手写数字
import os
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
from torchvision.utils import save_image

# ---------------- 超参数 ----------------
T = 1000 # 扩散步数
BETA_MIN = 1e-4 # 噪声调度最小值
BETA_MAX = 0.02 # 噪声调度最大值
BATCH_SIZE = 128
EPOCHS = 10
LR = 2e-4
TIME_DIM = 128 # 时间嵌入维度
BASE_CH = 32 # U-Net 基础通道数
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
SAVE_DIR = "output/basic_ddpm"
os.makedirs(SAVE_DIR, exist_ok=True)


# ---------------- 噪声调度系数 ----------------
def make_schedule(T, beta_min, beta_max, device):
"""预计算 DDPM 所需的所有系数"""
betas = torch.linspace(beta_min, beta_max, T, device=device) # β_t
alphas = 1.0 - betas # α_t = 1 - β_t
alphas_cumprod = torch.cumprod(alphas, dim=0) # ᾱ_t
alphas_cumprod_prev = F.pad(alphas_cumprod[:-1], (1, 0), value=1.0)

return {
"betas": betas,
"alphas": alphas,
"alphas_cumprod": alphas_cumprod,
"sqrt_alphas_cumprod": torch.sqrt(alphas_cumprod), # √ᾱ_t
"sqrt_one_minus_alphas_cumprod": torch.sqrt(1.0 - alphas_cumprod),# √(1-ᾱ_t)
"sqrt_recip_alphas": torch.sqrt(1.0 / alphas), # 1/√α_t
"posterior_variance": betas * (1.0 - alphas_cumprod_prev) / (1.0 - alphas_cumprod), # β̃_t
}

COEF = make_schedule(T, BETA_MIN, BETA_MAX, DEVICE)


def extract(a, t, shape):
"""从系数张量 a 中取出 t 时刻的值,并 reshape 成可与图像广播的形状"""
out = a.gather(0, t)
return out.reshape(t.shape[0], *((1,) * (len(shape) - 1)))


# ---------------- 前向加噪 ----------------
def q_sample(x0, t, noise=None):
"""一步加噪: x_t = √ᾱ_t * x0 + √(1-ᾱ_t) * ε"""
if noise is None:
noise = torch.randn_like(x0)
return (extract(COEF["sqrt_alphas_cumprod"], t, x0.shape) * x0
+ extract(COEF["sqrt_one_minus_alphas_cumprod"], t, x0.shape) * noise)


# ---------------- 时间嵌入 ----------------
class SinusoidalTimeEmbedding(nn.Module):
"""正弦时间嵌入,与 Transformer 的位置编码相同"""
def __init__(self, dim):
super().__init__()
self.dim = dim

def forward(self, t):
half = self.dim // 2
freqs = torch.exp(-math.log(10000) * torch.arange(half, device=t.device) / half)
args = t.float()[:, None] * freqs[None, :]
return torch.cat([torch.sin(args), torch.cos(args)], dim=-1) # (B, dim)


# ---------------- U-Net 模块 ----------------
class ResBlock(nn.Module):
"""残差块: (GroupNorm + SiLU + Conv) x 2,中间注入时间嵌入,带 skip 连接"""
def __init__(self, in_ch, out_ch, time_dim):
super().__init__()
self.norm1 = nn.GroupNorm(8, in_ch)
self.conv1 = nn.Conv2d(in_ch, out_ch, 3, padding=1)
self.norm2 = nn.GroupNorm(8, out_ch)
self.conv2 = nn.Conv2d(out_ch, out_ch, 3, padding=1)
self.time_proj = nn.Linear(time_dim, out_ch)
# 通道数变化时用 1x1 卷积对齐 skip 分支
self.skip = nn.Conv2d(in_ch, out_ch, 1) if in_ch != out_ch else nn.Identity()

def forward(self, x, t_emb):
h = self.conv1(F.silu(self.norm1(x)))
# 时间嵌入投影到通道维,broadcast 加到每个像素上
h = h + self.time_proj(F.silu(t_emb))[:, :, None, None]
h = self.conv2(F.silu(self.norm2(h)))
return h + self.skip(x)


class UNet(nn.Module):
"""
简化版 U-Net(针对 28x28 单通道图像):
28x28 -> 14x14 -> 7x7 -> 14x14 -> 28x28
每层通过跳跃连接拼接编码器特征。
"""
def __init__(self, base_ch=32, time_dim=128):
super().__init__()
self.time_mlp = nn.Sequential(
SinusoidalTimeEmbedding(time_dim),
nn.Linear(time_dim, time_dim * 4),
nn.SiLU(),
nn.Linear(time_dim * 4, time_dim),
)

c1, c2, c3 = base_ch, base_ch * 2, base_ch * 4
# 输入卷积: 不带归一化,保留原始带噪图像的灰度/噪声幅度信息
self.in_conv = nn.Conv2d(1, c1, 3, padding=1)
# 编码器
self.enc2 = ResBlock(c1, c2, time_dim)
self.enc3 = ResBlock(c2, c3, time_dim)
# 中间层
self.mid1 = ResBlock(c3, c3, time_dim)
self.mid2 = ResBlock(c3, c3, time_dim)
# 解码器: 输入通道 = 上采样通道 + 跳跃连接通道
self.dec2 = ResBlock(c3 + c2, c2, time_dim) # cat(上采样mid, e2) @14x14
self.dec1 = ResBlock(c2 + c1, c1, time_dim) # cat(上采样d2, e1) @28x28
# 输出头: 先归一化再卷积,稳定输出噪声预测的尺度
self.out_norm = nn.GroupNorm(8, c1)
self.out = nn.Conv2d(c1, 1, 3, padding=1)

def forward(self, x, t):
t_emb = self.time_mlp(t)

e1 = self.in_conv(x) # (B, 32, 28, 28)
e2 = self.enc2(F.avg_pool2d(e1, 2), t_emb) # (B, 64, 14, 14)
e3 = self.enc3(F.avg_pool2d(e2, 2), t_emb) # (B, 128, 7, 7)

m = self.mid2(self.mid1(e3, t_emb), t_emb) # (B, 128, 7, 7)

d2 = F.interpolate(m, scale_factor=2, mode="nearest") # 7 -> 14
d2 = self.dec2(torch.cat([d2, e2], dim=1), t_emb) # (B, 64, 14, 14)
d1 = F.interpolate(d2, scale_factor=2, mode="nearest") # 14 -> 28
d1 = self.dec1(torch.cat([d1, e1], dim=1), t_emb) # (B, 32, 28, 28)

return self.out(F.silu(self.out_norm(d1))) # (B, 1, 28, 28),预测的噪声


# ---------------- 训练 ----------------
def train():
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,)), # 缩放到 [-1, 1]
])
dataset = datasets.MNIST(root="./data", train=True, download=True, transform=transform)
loader = DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=True,
num_workers=2, drop_last=True)

model = UNet(BASE_CH, TIME_DIM).to(DEVICE)
optimizer = torch.optim.AdamW(model.parameters(), lr=LR)

for epoch in range(EPOCHS):
total_loss = 0.0
for x0, _ in loader:
x0 = x0.to(DEVICE)
# 1. 随机时间步与噪声
t = torch.randint(0, T, (x0.shape[0],), device=DEVICE)
noise = torch.randn_like(x0)
# 2. 前向加噪
xt = q_sample(x0, t, noise)
# 3. 预测噪声并计算损失
pred = model(xt, t)
loss = F.mse_loss(pred, noise)

optimizer.zero_grad()
loss.backward()
optimizer.step()
total_loss += loss.item()

print(f"Epoch {epoch+1:03d}/{EPOCHS} loss={total_loss/len(loader):.4f}")
# 每个 epoch 保存一次模型并采样观察
torch.save(model.state_dict(), f"{SAVE_DIR}/ddpm_mnist.pt")
sample(model, n=16, fname=f"{SAVE_DIR}/epoch_{epoch+1:03d}.png")


# ---------------- 采样 ----------------
@torch.no_grad()
def sample(model, n=16, fname=None):
"""DDPM 反向采样: 从纯噪声逐步去噪"""
model.eval()
x = torch.randn(n, 1, 28, 28, device=DEVICE)

for i in reversed(range(T)):
t = torch.full((n,), i, device=DEVICE, dtype=torch.long)
# 网络预测噪声
eps = model(x, t)
# 由预测噪声计算均值 μ_θ
coef1 = extract(COEF["sqrt_recip_alphas"], t, x.shape)
coef2 = extract(COEF["betas"], t, x.shape) / extract(
COEF["sqrt_one_minus_alphas_cumprod"], t, x.shape)
mean = coef1 * (x - coef2 * eps)
# t > 0 时加入随机噪声,t = 0 时直接取均值
if i > 0:
var = extract(COEF["posterior_variance"], t, x.shape)
x = mean + torch.sqrt(var) * torch.randn_like(x)
else:
x = mean

x = (x.clamp(-1, 1) + 1) / 2 # 从 [-1,1] 还原到 [0,1]
if fname:
save_image(x, fname, nrow=int(n ** 0.5))
model.train()
return x


if __name__ == "__main__":
train()

运行方式:

1
python basic_ddpm.py

训练过程中会在 output/basic_ddpm/ 目录下保存模型权重 ddpm_mnist.pt 和每个 epoch 的采样结果 epoch_XXX.png。在单张 GPU 上约 10 个 epoch 即可看到清晰的手写数字。

训练完成后的生成结果(16 张随机样本):

DDPM 生成结果

6. DDIM:加速采样

6.1 动机

DDPM 的反向过程是一条马尔可夫链,必须严格走完 t=T,T−1,…,1t=T, T-1, \dots, 1 共 TT 步,每步都要跑一次网络前向。T=1000T=1000 时生成一张图需要 1000 次推理,非常慢。

DDIM(Denoising Diffusion Implicit Models)的核心观察是:DDPM 的训练目标只依赖边缘分布 q(xt∣x0)q(x_t\mid x_0),而不依赖前向过程的联合分布。也就是说,可以构造一族不同的前向过程,它们拥有完全相同的 q(xt∣x0)q(x_t\mid x_0),因此训练好的 DDPM 模型无需重新训练即可直接使用,但其中一些前向过程对应的反向采样可以跳步、甚至完全确定性地进行。

6.2 非马尔可夫前向过程

DDIM 定义前向过程为(σt\sigma_t 是可自由选择的参数):

qσ(xt−1∣xt,x0)=N(xt−1; αˉt−1 x0+1−αˉt−1−σt2⋅xt−αˉt x01−αˉt, σt2I)q_\sigma(x_{t-1}\mid x_t, x_0) = \mathcal{N}\left(x_{t-1};\ \sqrt{\bar{\alpha}_{t-1}}\,x_0 + \sqrt{1-\bar{\alpha}_{t-1}-\sigma_t^2}\cdot\frac{x_t-\sqrt{\bar{\alpha}_t}\,x_0}{\sqrt{1-\bar{\alpha}_t}},\ \sigma_t^2\mathbf{I}\right)

注意它直接以 x0x_0 为条件(非马尔可夫),但边缘分布仍然是 q(xt∣x0)=N(αˉtx0,(1−αˉt)I)q(x_t\mid x_0)=\mathcal{N}(\sqrt{\bar{\alpha}_t}x_0,(1-\bar{\alpha}_t)\mathbf{I}),与 DDPM 一致,所以训练目标完全不变。

两个重要的特例:

  • 当 σt=1−αˉt−11−αˉt1−αˉtαˉt−1\sigma_t = \sqrt{\frac{1-\bar{\alpha}_{t-1}}{1-\bar{\alpha}_t}}\sqrt{1-\frac{\bar{\alpha}_t}{\bar{\alpha}_{t-1}}} 时,退化为 DDPM 的马尔可夫前向过程;
  • 当 σt=0\sigma_t = 0 时,前向与反向过程都变成确定性的——这就是"implicit model"名字的由来。

实际使用中通常引入系数 η∈[0,1]\eta\in[0,1] 在两者之间插值:

σt(η)=η1−αˉt−11−αˉt1−αˉtαˉt−1\sigma_t(\eta) = \eta\sqrt{\frac{1-\bar{\alpha}_{t-1}}{1-\bar{\alpha}_t}}\sqrt{1-\frac{\bar{\alpha}_t}{\bar{\alpha}_{t-1}}}

η=0\eta=0 为确定性采样,η=1\eta=1 等价于 DDPM。

6.3 DDIM 采样公式

将网络预测的噪声 ϵθ(xt,t)\epsilon_\theta(x_t,t) 代入,先估计"当前噪声图对应的原始图像":

x^0=xt−1−αˉt ϵθ(xt,t)αˉt\hat{x}_0 = \frac{x_t - \sqrt{1-\bar{\alpha}_t}\,\epsilon_\theta(x_t,t)}{\sqrt{\bar{\alpha}_t}}

然后一步得到 xt−1x_{t-1}:

xt−1=αˉt−1 x^0+1−αˉt−1−σt2 ϵθ(xt,t)+σtz,z∼N(0,I)\boxed{x_{t-1} = \sqrt{\bar{\alpha}_{t-1}}\,\hat{x}_0 + \sqrt{1-\bar{\alpha}_{t-1}-\sigma_t^2}\,\epsilon_\theta(x_t,t) + \sigma_t z,\qquad z\sim\mathcal{N}(0,\mathbf{I})}

直观理解:x^0\hat{x}_0 是"对最终结果的猜测",公式把它与"指向当前 xtx_t 的方向(噪声)"按新的噪声水平重新组合。

6.4 跳步采样

因为训练目标对任意 tt 都成立,采样时不必使用全部 TT 步。任选一个长度 S≪TS \ll T 的时间步子序列

{τ1<τ2<⋯<τS}⊂{1,…,T}\{\tau_1 < \tau_2 < \dots < \tau_S\} \subset \{1,\dots,T\}

只在子序列上执行上述更新公式(公式中的 αˉt−1\bar{\alpha}_{t-1} 换成 αˉτi−1\bar{\alpha}_{\tau_{i-1}}),即可用 SS 次网络推理完成生成。典型取 S=50S=50,速度提升 20 倍,而图像质量几乎无损。

DDIM 还带来两个额外好处:

  1. 确定性生成(η=0\eta=0):相同的初始噪声 xTx_T 总是生成相同的图像;
  2. 潜空间插值:两个初始噪声之间做球面插值,生成结果会平滑过渡,说明 xTx_T 是一个语义有意义的隐变量。

6.5 实现

basic_ddim.py 是一个可独立运行的完整文件:模型结构(ResBlock/UNet)、噪声调度(make_schedule/q_sample)和训练循环与第 5 节的 basic_ddpm.py 完全相同(DDIM 复用 DDPM 的训练目标,无需改动训练代码),相同部分不再重复列出。区别只有三处:噪声调度输出目录为 output/basic_ddim、权重保存为 ddim_mnist.pt、每个 epoch 用下面的 ddim_sample 代替 DDPM 的 sample——该函数与 DDPM 共用同一套调度系数 COEF,无需重新训练即可直接加载 DDPM 训练好的权重:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
@torch.no_grad()
def ddim_sample(model, n=16, steps=50, eta=0.0, fname=None):
"""
DDIM 采样: 跳步 + 可选确定性
steps: 实际采样步数 S(远小于 T)
eta: 0.0 为确定性采样, 1.0 等价于 DDPM 的随机采样
"""
model.eval()
# 构造长度 steps 的时间步子序列(等间隔,含首尾)
times = torch.linspace(0, T - 1, steps, device=DEVICE).long()

x = torch.randn(n, 1, 28, 28, device=DEVICE)
for i in reversed(range(len(times))):
t = times[i]
# 前一个时间步的 ᾱ;i=0 时对应 ᾱ_0 = 1(即恢复 x_0)
alpha_t = COEF["alphas_cumprod"][t]
alpha_prev = COEF["alphas_cumprod"][times[i - 1]] if i > 0 else torch.tensor(1.0, device=DEVICE)

t_batch = torch.full((n,), t.item(), device=DEVICE, dtype=torch.long)
eps = model(x, t_batch)

# 1. 估计原始图像 x̂0
x0_pred = (x - torch.sqrt(1.0 - alpha_t) * eps) / torch.sqrt(alpha_t)

# 2. 计算 σ_t(η)
sigma = eta * torch.sqrt((1.0 - alpha_prev) / (1.0 - alpha_t)) \
* torch.sqrt(1.0 - alpha_t / alpha_prev)

# 3. 组合得到 x_{t-1}
dir_xt = torch.sqrt((1.0 - alpha_prev - sigma ** 2).clamp(min=0)) * eps
x = torch.sqrt(alpha_prev) * x0_pred + dir_xt
if i > 0 and sigma > 0:
x = x + sigma * torch.randn_like(x)

x = (x.clamp(-1, 1) + 1) / 2
if fname:
save_image(x, fname, nrow=int(n ** 0.5))
model.train()
return x

train() 中保存权重与采样的两行相应改为(其余训练代码与 basic_ddpm.py 一致):

1
2
torch.save(model.state_dict(), f"{SAVE_DIR}/ddim_mnist.pt")
ddim_sample(model, n=16, steps=50, eta=0.0, fname=f"{SAVE_DIR}/epoch_{epoch+1:03d}.png")

运行方式:

1
python basic_ddim.py

输出保存在 output/basic_ddim/。与 DDPM 对比可以直观观察速度差异:

对比实验建议:

配置 推理次数 特点
DDPM(第 5 节 sample) 1000 原始质量,最慢
ddim_sample(steps=50, eta=0) 50 确定性,质量接近 DDPM
ddim_sample(steps=20, eta=0) 20 更快,开始出现细节损失
ddim_sample(steps=50, eta=1) 50 随机采样,多样性略增

确定性(η=0\eta=0)可以验证:固定随机种子两次调用 ddim_sample,生成的图像完全一致;而 DDPM 采样每次结果都不同。

DDIM 50 步确定性采样(η=0\eta=0)的生成结果:

DDIM 生成结果

7. 代码与公式的对应关系(DDPM/DDIM)

公式 代码位置
βt\beta_t 线性调度 make_schedule 中的 torch.linspace
xt=αˉtx0+1−αˉtϵx_t = \sqrt{\bar\alpha_t}x_0 + \sqrt{1-\bar\alpha_t}\epsilon q_sample()
正弦时间嵌入 SinusoidalTimeEmbedding
噪声预测网络 ϵθ(xt,t)\epsilon_\theta(x_t,t) UNet
L=∣ϵ−ϵθ(xt,t)∣2\mathcal{L} = |\epsilon - \epsilon_\theta(x_t,t)|^2 train() 中的 F.mse_loss(pred, noise)
μθ=1αt(xt−βt1−αˉtϵθ)\mu_\theta = \frac{1}{\sqrt{\alpha_t}}\big(x_t - \frac{\beta_t}{\sqrt{1-\bar\alpha_t}}\epsilon_\theta\big) sample() 中的 mean
xt−1=μθ+β~t zx_{t-1} = \mu_\theta + \sqrt{\tilde\beta_t}\,z sample() 循环体
x^0=(xt−1−αˉtϵθ)/αˉt\hat{x}_0 = \big(x_t-\sqrt{1-\bar\alpha_t}\epsilon_\theta\big)/\sqrt{\bar\alpha_t} ddim_sample() 中的 x0_pred
DDIM 更新 xt−1=αˉt−1x^0+1−αˉt−1−σ2ϵθ+σzx_{t-1}=\sqrt{\bar\alpha_{t-1}}\hat{x}_0+\sqrt{1-\bar\alpha_{t-1}-\sigma^2}\epsilon_\theta+\sigma z ddim_sample() 循环体

第二部分:Mini Stable Diffusion

基础 DDPM 已经能在像素空间生成图像,但距离 Stable Diffusion(SD)还差三件事:
潜空间扩散、条件注入、Classifier-Free Guidance。本部分在第一部分的基础上逐一加入。

8. 从 DDPM 到 Stable Diffusion:三个增量

增量 解决的问题 SD 中的部件 本实现中的对应
① 潜空间扩散 像素空间计算量大 KL-VAE tiny AutoEncoder
② 条件注入 无法控制生成内容 CLIP 文本编码器 + cross-attention 类别标签 embedding(“文本” = 数字 0~9)
③ Classifier-Free Guidance 条件强度不可调 CFG CFG

整体架构(LDM / SD 的标准结构):

1
2
3
4
5
6
7
8
9
       像素空间                          潜空间
┌──────────────────────────────────────┐
x ──enc──► z_0 ──► 前向加噪 z_0 → z_1 → ... → z_T │
(图像) │ │ 扩散过程全部
│ 反向去噪: ε_θ(z_t, t, c) U-Net │ 在潜空间进行
x̂ ◄──dec── z_0 ◄── z_T → ... → z_1 → z_0 │
└──────────────────────────────────────┘
▲
条件 c(文本/标签)经编码后注入 U-Net 每一层

训练好之后的生成流程:先在潜空间中从纯噪声去噪得到 z0z_0,再用解码器一次性还原为图像。

9. 增量①:潜空间扩散(Latent Diffusion)

9.1 为什么不在像素空间做扩散

一张 512×512×3512\times512\times3 的图像有 786432 维,但其中绝大部分维度承载的是人眼几乎不可察觉的高频细节。DDPM 的每一步去噪都要在这个巨大的维度上跑一次 U-Net,训练和采样都非常昂贵。

LDM(Latent Diffusion Model,SD 的正式名称)的思路:先用一个自编码器把图像压缩到一个语义等效但维度小得多的潜空间,扩散过程全部在潜空间进行。SD 的 VAE 把 512×512×3512\times512\times3 压到 64×64×464\times64\times4(48 倍压缩),计算量随之大幅下降,而生成质量几乎无损。

9.2 自编码器与潜变量缩放

SD 使用的是 KL-VAE(编码器输出高斯分布参数,用 KL 散度正则化潜空间)。在 mini 版本中用普通自编码器(MSE 重建损失)即可达到同样目的:

LAE=∥x−dec(enc(x))∥2\mathcal{L}_{AE} = \|x - \text{dec}(\text{enc}(x))\|^2

AE 先单独训练若干轮,然后冻结,作为像素空间与潜空间之间固定的桥梁。

一个重要的工程细节:AE 编码出的潜变量 zz 的尺度是任意的,而扩散过程的噪声调度 αˉt\bar{\alpha}_t 是为标准差约为 1 的数据设计的。因此需要对潜变量做归一化:

znorm=zσz,σz=训练集上潜变量的标准差z_{\text{norm}} = \frac{z}{\sigma_z},\qquad \sigma_z = \text{训练集上潜变量的标准差}

解码时再乘回 σz\sigma_z。SD 中那个著名的缩放因子 0.18215 起的就是这个作用(它是 SD 的 VAE 在训练数据上潜变量标准差的倒数)。

9.3 训练与采样的变化

与像素空间 DDPM 相比,公式完全不变,只是把所有 xx 换成 zz:

  • 训练:z0=enc(x)/σzz_0 = \text{enc}(x)/\sigma_z,然后照常加噪 zt=αˉtz0+1−αˉtϵz_t = \sqrt{\bar{\alpha}_t}z_0 + \sqrt{1-\bar{\alpha}_t}\epsilon,预测噪声;
  • 采样:DDIM 在潜空间走完得到 z0z_0,最后 x^=dec(z0⋅σz)\hat{x} = \text{dec}(z_0 \cdot \sigma_z)。

10. 增量②:条件注入

10.1 真实 SD 的做法

SD 的条件是自由文本。CLIP 文本编码器把 prompt 编码为 77 个 token 的特征序列,U-Net 的每个 ResBlock 后面接一个 cross-attention 层:图像特征作为 query,文本 token 特征作为 key/value,让每个图像区域"读取"相关的文本信息。

10.2 mini 版的简化

本实现用数字类别 0~9 作为"文本"。此时条件序列只有 1 个 token——而 cross-attention 对单个 key 的 softmax 恒为 1,attention 退化成一次线性变换再加法。所以 mini 版可以直接采用与时间步 tt 相同的注入方式:

h←h+Wt ϕ(t)+Wc ψ(c)h \leftarrow h + W_t\,\phi(t) + W_c\,\psi(c)

其中 ψ(c)\psi(c) 是可学习的类别 embedding(等价于一个 mini 版"文本编码器"),WcW_c 把它投影到特征通道数后逐像素相加。结构上与真实 SD 完全同构,将来换成多 token 条件时,把"投影相加"换成 cross-attention 即可。

11. 增量③:Classifier-Free Guidance(CFG)

11.1 动机与原理

朴素条件生成中,条件 cc 对结果的约束力往往不够强——模型容易"忽视"条件。我们希望放大条件的影响。

由贝叶斯公式 p(c∣x)∝p(x∣c)/p(x)p(c\mid x) \propto p(x\mid c)/p(x),对 score(对数概率梯度)有:

∇xlog⁡p(x∣c)=∇xlog⁡p(x)+∇xlog⁡p(c∣x)\nabla_x \log p(x\mid c) = \nabla_x \log p(x) + \nabla_x \log p(c\mid x)

即条件 score = 无条件 score + 一个隐式分类器的梯度。给分类器项加个权重 ww,就得到加强版的条件 score。由于扩散模型预测的是噪声而噪声与 score 只相差一个系数,可以直接写成噪声预测的外推形式:

ϵ~θ(zt,t,c)=ϵθ(zt,t,∅)+w [ϵθ(zt,t,c)−ϵθ(zt,t,∅)]\boxed{\tilde{\epsilon}_\theta(z_t, t, c) = \epsilon_\theta(z_t, t, \varnothing) + w\,\big[\epsilon_\theta(z_t, t, c) - \epsilon_\theta(z_t, t, \varnothing)\big]}

  • w=0w=0:完全无条件生成
  • w=1w=1:普通条件生成
  • w>1w>1:条件约束力增强(SD 常用 7.5 左右;过大会过饱和、出现伪影)

11.2 训练:条件丢弃

注意公式里需要同一个网络同时能输出条件预测和无条件预测,这通过训练时的"条件丢弃"实现:以一定概率(通常 10%)把条件 cc 替换为一个特殊的空条件 ∅\varnothing(本实现中为第 11 个类别 embedding)。这样网络同时学会了 p(x∣c)p(x\mid c) 和 p(x)p(x)。

11.3 采样:双前向外推

采样时每个时间步跑两次网络前向(条件、无条件各一次),按上面的公式外推,再用外推后的 ϵ~\tilde{\epsilon} 走 DDIM 更新。代价是推理计算量翻倍。

12. Mini SD 的 MNIST 完整实现

单文件实现:tiny AE(28×28×1↔14×14×428\times28\times1 \leftrightarrow 14\times14\times4)+
条件 U-Net + DDIM/CFG 采样器。环境要求:torch、torchvision。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
# mini_sd_mnist.py
# 从零实现的 Mini Stable Diffusion:
# ① 潜空间扩散(tiny AE)② 类别条件注入 ③ Classifier-Free Guidance
import os
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
from torchvision.utils import save_image

# ---------------- 超参数 ----------------
T = 1000 # 扩散步数
BETA_MIN, BETA_MAX = 1e-4, 0.02
BATCH_SIZE = 128
AE_EPOCHS = 2 # 自编码器预训练轮数
DIFF_EPOCHS = 10 # 扩散模型训练轮数
AE_LR, DIFF_LR = 1e-3, 2e-4
LATENT_CH = 4 # 潜空间通道数(28x28x1 -> 14x14x4)
BASE_CH = 64
EMB_DIM = 128 # 时间/条件 embedding 维度
NUM_CLASSES = 10
NULL_CLASS = NUM_CLASSES # CFG 的空条件(第 11 个 embedding)
CFG_DROP_PROB = 0.1 # 训练时条件丢弃概率
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
SAVE_DIR = "output/mini_sd"
os.makedirs(SAVE_DIR, exist_ok=True)


# ---------------- 噪声调度(与 DDPM 完全相同)----------------
betas = torch.linspace(BETA_MIN, BETA_MAX, T, device=DEVICE)
alphas_cumprod = torch.cumprod(1.0 - betas, dim=0) # ᾱ_t


def extract(a, t, shape):
"""取出 a[t] 并 reshape 成可广播到图像的形状"""
return a.gather(0, t).reshape(t.shape[0], *((1,) * (len(shape) - 1)))


# ---------------- ① 自编码器(像素 <-> 潜空间)----------------
class TinyAutoencoder(nn.Module):
"""28x28x1 <-> 14x14x4,MSE 重建训练,代替 SD 的 KL-VAE"""
def __init__(self, latent_ch=4):
super().__init__()
self.encoder = nn.Sequential(
nn.Conv2d(1, 32, 3, stride=2, padding=1), # 28 -> 14
nn.SiLU(),
nn.Conv2d(32, latent_ch, 3, padding=1),
)
self.decoder = nn.Sequential(
nn.Conv2d(latent_ch, 32, 3, padding=1),
nn.SiLU(),
nn.ConvTranspose2d(32, 1, 4, stride=2, padding=1), # 14 -> 28
nn.Tanh(), # 输出范围 [-1,1],与数据归一化一致
)

def forward(self, x):
return self.decoder(self.encoder(x))


# ---------------- 时间/条件嵌入 ----------------
class SinusoidalTimeEmbedding(nn.Module):
def __init__(self, dim):
super().__init__()
self.dim = dim

def forward(self, t):
half = self.dim // 2
freqs = torch.exp(-math.log(10000) * torch.arange(half, device=t.device) / half)
args = t.float()[:, None] * freqs[None, :]
return torch.cat([torch.sin(args), torch.cos(args)], dim=-1)


# ---------------- ② 条件 U-Net ----------------
class CondResBlock(nn.Module):
"""ResBlock + 双路条件注入:时间 t 与条件 c 用同样方式(投影到通道后相加)"""
def __init__(self, in_ch, out_ch, emb_dim):
super().__init__()
self.norm1 = nn.GroupNorm(min(8, in_ch), in_ch)
self.conv1 = nn.Conv2d(in_ch, out_ch, 3, padding=1)
self.norm2 = nn.GroupNorm(min(8, out_ch), out_ch)
self.conv2 = nn.Conv2d(out_ch, out_ch, 3, padding=1)
self.t_proj = nn.Linear(emb_dim, out_ch) # 时间条件投影
self.c_proj = nn.Linear(emb_dim, out_ch) # 类别条件投影
self.skip = nn.Conv2d(in_ch, out_ch, 1) if in_ch != out_ch else nn.Identity()

def forward(self, x, t_emb, c_emb):
h = self.conv1(F.silu(self.norm1(x)))
h = h + self.t_proj(F.silu(t_emb))[:, :, None, None]
h = h + self.c_proj(F.silu(c_emb))[:, :, None, None]
h = self.conv2(F.silu(self.norm2(h)))
return h + self.skip(x)


class CondUNet(nn.Module):
"""在潜空间(14x14x4)工作的条件 U-Net,预测噪声 ε_θ(z_t, t, c)"""
def __init__(self, latent_ch=4, base_ch=64, emb_dim=128):
super().__init__()
self.time_mlp = nn.Sequential(
SinusoidalTimeEmbedding(emb_dim),
nn.Linear(emb_dim, emb_dim), nn.SiLU(),
nn.Linear(emb_dim, emb_dim),
)
# mini 版"文本编码器":类别 -> 向量,多一类 NULL_CLASS 给 CFG 空条件
self.label_emb = nn.Embedding(NUM_CLASSES + 1, emb_dim)

c1, c2 = base_ch, base_ch * 2
self.in_conv = nn.Conv2d(latent_ch, c1, 3, padding=1)
self.down = CondResBlock(c1, c2, emb_dim) # 14x14 -> 7x7
self.mid1 = CondResBlock(c2, c2, emb_dim)
self.mid2 = CondResBlock(c2, c2, emb_dim)
self.up = CondResBlock(c2 + c1, c1, emb_dim) # 7x7 -> 14x14(拼接跳跃连接)
self.out_norm = nn.GroupNorm(8, c1)
self.out_conv = nn.Conv2d(c1, latent_ch, 3, padding=1)

def forward(self, z, t, y):
t_emb = self.time_mlp(t)
c_emb = self.label_emb(y)
h0 = self.in_conv(z) # (B, 64, 14, 14)
h1 = self.down(F.avg_pool2d(h0, 2), t_emb, c_emb) # (B, 128, 7, 7)
m = self.mid2(self.mid1(h1, t_emb, c_emb), t_emb, c_emb)
u = torch.cat([F.interpolate(m, scale_factor=2, mode="nearest"), h0], dim=1)
u = self.up(u, t_emb, c_emb) # (B, 64, 14, 14)
return self.out_conv(F.silu(self.out_norm(u))) # (B, 4, 14, 14)


# ---------------- 训练 ----------------
def get_dataloader():
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,)), # [-1, 1]
])
dataset = datasets.MNIST(root="./data", train=True, download=True, transform=transform)
return DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=True,
num_workers=2, drop_last=True)


def train():
loader = get_dataloader()
ae = TinyAutoencoder(LATENT_CH).to(DEVICE)
model = CondUNet(LATENT_CH, BASE_CH, EMB_DIM).to(DEVICE)

# ---- 阶段一:训练自编码器(像素 <-> 潜空间的桥梁)----
opt_ae = torch.optim.AdamW(ae.parameters(), lr=AE_LR)
for epoch in range(AE_EPOCHS):
total = 0.0
for x, _ in loader:
x = x.to(DEVICE)
loss = F.mse_loss(ae(x), x) # 重建损失
opt_ae.zero_grad(); loss.backward(); opt_ae.step()
total += loss.item()
print(f"[AE] epoch {epoch+1}/{AE_EPOCHS} recon_loss={total/len(loader):.5f}")

# 冻结 AE,并估计潜变量标准差(SD 中 0.18215 缩放因子的同款作用)
ae.requires_grad_(False).eval()
with torch.no_grad():
x, _ = next(iter(loader))
latent_std = ae.encoder(x.to(DEVICE)).std().item()
print(f"latent_std = {latent_std:.3f}")

# ---- 阶段二:在潜空间训练条件扩散模型 ----
opt = torch.optim.AdamW(model.parameters(), lr=DIFF_LR)
for epoch in range(DIFF_EPOCHS):
total = 0.0
for x, y in loader:
x, y = x.to(DEVICE), y.to(DEVICE)
with torch.no_grad():
z0 = ae.encoder(x) / latent_std # 增量①:在潜空间扩散
# 增量③:以 10% 概率把条件换成空条件(CFG 训练)
drop = torch.rand(y.shape[0], device=DEVICE) < CFG_DROP_PROB
y = torch.where(drop, torch.full_like(y, NULL_CLASS), y)

t = torch.randint(0, T, (z0.shape[0],), device=DEVICE)
eps = torch.randn_like(z0)
zt = (extract(alphas_cumprod.sqrt(), t, z0.shape) * z0
+ extract((1 - alphas_cumprod).sqrt(), t, z0.shape) * eps)
pred = model(zt, t, y) # 增量②:多接一个条件 y
loss = F.mse_loss(pred, eps)

opt.zero_grad(); loss.backward(); opt.step()
total += loss.item()
print(f"[Diff] epoch {epoch+1}/{DIFF_EPOCHS} loss={total/len(loader):.4f}")

torch.save({"model": model.state_dict(), "ae": ae.state_dict(),
"latent_std": latent_std}, f"{SAVE_DIR}/mini_sd.pt")
# 每个 epoch 生成 0~9 各一张(w=3.0)观察效果
sample(model, ae, latent_std, torch.arange(10, device=DEVICE),
steps=50, w=3.0, fname=f"{SAVE_DIR}/epoch_{epoch+1:03d}.png")

# ---------------- ③ 采样:DDIM + CFG ----------------
@torch.no_grad()
def sample(model, ae, latent_std, y, steps=50, eta=0.0, w=3.0, fname=None):
"""
在潜空间做 DDIM 采样,每步用 CFG 双前向外推,最后解码回像素。
y: 想生成的类别(batch 内可以各不相同)
steps: DDIM 跳步数
w: CFG 引导强度,w=0 无条件,w=1 普通条件,w>1 强条件
"""
model.eval()
B = y.shape[0]
y_null = torch.full_like(y, NULL_CLASS) # 无条件分支
times = torch.linspace(T - 1, 0, steps, device=DEVICE).long()
z = torch.randn(B, LATENT_CH, 14, 14, device=DEVICE)

for i in range(len(times)):
t = times[i]
alpha_t = alphas_cumprod[t]
alpha_prev = alphas_cumprod[times[i + 1]] if i + 1 < len(times) \
else torch.tensor(1.0, device=DEVICE)
t_batch = torch.full((B,), t.item(), device=DEVICE, dtype=torch.long)

# CFG:条件与无条件各前向一次,然后外推
eps_cond = model(z, t_batch, y)
eps_uncond = model(z, t_batch, y_null)
eps = eps_uncond + w * (eps_cond - eps_uncond)

# DDIM 更新(与第一部分的公式一致)
z0_pred = (z - torch.sqrt(1 - alpha_t) * eps) / torch.sqrt(alpha_t)
sigma = eta * torch.sqrt((1 - alpha_prev) / (1 - alpha_t)) \
* torch.sqrt(1 - alpha_t / alpha_prev)
dir_zt = torch.sqrt((1 - alpha_prev - sigma ** 2).clamp(min=0)) * eps
z = torch.sqrt(alpha_prev) * z0_pred + dir_zt
if i + 1 < len(times) and sigma > 0:
z = z + sigma * torch.randn_like(z)

x = ae.decoder(z * latent_std) # 潜空间 -> 像素空间
x = (x.clamp(-1, 1) + 1) / 2
if fname:
save_image(x, fname, nrow=B)
model.train()
return x


if __name__ == "__main__":
train()

运行方式:

1
python mini_sd.py

训练分两个阶段:先训 2 个 epoch 的自编码器并冻结,再在潜空间训 10 个 epoch 的扩散模型。输出保存在 output/mini_sd/,每个 epoch 生成一张 0~9 的数字条带图(nrow=10,每列一个类别),可以直观看到条件控制是否生效。

训练完成后的条件生成结果(0~9 每类各一张,w=3.0w=3.0):

Mini SD 生成结果

13. 代码与公式的对应关系(Mini SD)

公式 代码位置
LAE=∣x−dec(enc(x))∣2\mathcal{L}_{AE} = |x - \text{dec}(\text{enc}(x))|^2 train() 阶段一中的 F.mse_loss(ae(x), x)
znorm=z/σzz_{\text{norm}} = z/\sigma_z(0.18215 同款) train() 中的 ae.encoder(x) / latent_std 与 sample() 末尾的 ae.decoder(z * latent_std)
zt=αˉtz0+1−αˉtϵz_t = \sqrt{\bar\alpha_t}z_0 + \sqrt{1-\bar\alpha_t}\epsilon train() 阶段二中的 zt
条件注入 h←h+Wtϕ(t)+Wcψ(c)h \leftarrow h + W_t\phi(t) + W_c\psi(c) CondResBlock.forward() 中的两次相加
CFG 条件丢弃 torch.where(drop, NULL_CLASS, y)
ϵ~=ϵ∅+w(ϵc−ϵ∅)\tilde{\epsilon} = \epsilon_{\varnothing} + w(\epsilon_c - \epsilon_{\varnothing}) sample() 中的 eps_uncond + w * (eps_cond - eps_uncond)
DDIM 更新 sample() 循环体
x^=dec(z0⋅σz)\hat{x} = \text{dec}(z_0\cdot\sigma_z) sample() 末尾的 ae.decoder(z * latent_std)

第三部分:Mini FLUX

FLUX(以及 Stable Diffusion 3)在 SD 的基础上又做了三处核心改动:
训练目标、采样器、主干网络。本部分在第二部分的基础上继续演进。

14. 从 SD 到 FLUX:换掉三件,继承三件

FLUX(以及 SD3)与经典 SD 的关系可以概括为"三换三不换":

SD(第二部分) FLUX(本部分)
潜空间(VAE/AE) ✅ 继承 ✅ 原样保留
条件机制 + CFG ✅ 继承 ✅ 公式一字不差
文本编码器 CLIP T5(本实现仍用类别 embedding)
训练目标 预测噪声 ϵ\epsilon ❌ 换成预测速度 vv(Flow Matching)
采样器 DDPM/DDIM(1000/50 步) ❌ 换成 Euler 解 ODE(约 20 步)
主干网络 U-Net ❌ 换成双流 MMDiT(纯 Transformer)

骨架不变,换掉的是三件内核:

1
2
3
4
5
6
7
不变:  像素 --AE--> 潜空间 --扩散--> 采样 --AE--> 像素
条件注入 + CFG

换掉: 加噪方式 z_t = √ᾱ·z0 + √(1-ᾱ)·ε → z_t = (1-t)·z0 + t·ε (直线插值)
预测目标 噪声 ε → 速度 v = ε - z0
采样方式 DDIM 逐步去噪 → Euler 解 ODE,dz/dt = v
主干网络 卷积 U-Net → 双流 MMDiT(Transformer)

15. Flow Matching:把扩散拉成直线

15.1 直线插值路径

回顾 DDPM 的加噪公式 zt=αˉtz0+1−αˉtϵz_t = \sqrt{\bar\alpha_t}z_0 + \sqrt{1-\bar\alpha_t}\epsilon:
z0z_0 与噪声 ϵ\epsilon 的混合系数随 tt 沿一条曲线变化,而且定义在 1000 个离散时间步上。

Flow Matching(本文档使用其最简洁的形式,即 rectified flow / SD3 与 FLUX 采用的形式化)把路径直接定义为直线插值,时间也改为连续的 t∈[0,1]t\in[0,1]:

zt=(1−t) z0+t ϵ,ϵ∼N(0,I)z_t = (1-t)\,z_0 + t\,\epsilon,\qquad \epsilon\sim\mathcal{N}(0,\mathbf{I})

t=0t=0 是数据,t=1t=1 是纯噪声。注意沿这条路径,ztz_t 关于 tt 的导数是常数:

dztdt=ϵ−z0≜v\frac{\mathrm{d}z_t}{\mathrm{d}t} = \epsilon - z_0 \triangleq v

即样本从数据流向噪声的速度。

15.2 训练目标:预测速度

既然速度就是 ϵ−z0\epsilon - z_0,训练一个网络 vθ(zt,t,c)v_\theta(z_t, t, c) 去回归它即可:

LFM=Ez0, t, ϵ∥vθ((1−t)z0+tϵ, t, c)−(ϵ−z0)∥2\boxed{\mathcal{L}_{FM} = \mathbb{E}_{z_0,\,t,\,\epsilon}\left\|v_\theta\big((1-t)z_0 + t\epsilon,\ t,\ c\big) - (\epsilon - z_0)\right\|^2}

与 DDPM 的训练循环对比,只有三处改动:

  1. tt 从离散均匀采样换成连续采样(见 15.4);
  2. 加噪公式换成直线插值;
  3. 回归目标从 ϵ\epsilon 换成 v=ϵ−z0v = \epsilon - z_0。

15.3 采样:Euler 解 ODE

vθv_\theta 逼近的是一条确定性的速度场,因此生成过程就是解常微分方程

dzdt=vθ(z,t,c)\frac{\mathrm{d}z}{\mathrm{d}t} = v_\theta(z, t, c)

从 t=1t=1(纯噪声)积分回 t=0t=0。最简单的 Euler 法:

z←z−vθ(z,t,c)⋅Δt,Δt=1Sz \leftarrow z - v_\theta(z, t, c)\cdot\Delta t,\qquad \Delta t = \frac{1}{S}

走 SS 步即完成生成。因为真实路径是直线,速度近乎常数,Euler 法的离散误差很小——这就是为什么 FLUX 用 20 步就能生成,而 DDPM 需要 1000 步。CFG 的用法与 SD 完全相同:每步对条件和无条件各前向一次,外推 v~=v∅+w(vc−v∅)\tilde v = v_\varnothing + w(v_c - v_\varnothing) 后再走 Euler 步。

15.4 时间步采样:logit-normal

训练时 tt 怎么采样是个重要细节。均匀采样会把大量训练预算浪费在接近两端(几乎无噪或几乎全噪)的"简单"时刻上。SD3/FLUX 采用 logit-normal 采样:

t=sigmoid(u),u∼N(0,1)t = \text{sigmoid}(u),\qquad u\sim\mathcal{N}(0,1)

sigmoid 把钟形分布压到 (0,1)(0,1) 区间,密度集中在 t=0.5t=0.5 附近——即"半噪半数据"、学习难度最高的中间时刻。

16. MMDiT:纯 Transformer 主干

FLUX/SD3 用 MMDiT(Multimodal Diffusion Transformer)取代 U-Net。它由三部分组成:patchify、NN 个双流 block、final layer。

16.1 Patchify 与位置编码

Transformer 处理的是 token 序列,所以先把潜变量切成 patch。14×14×414\times14\times4 的潜变量用 2×22\times2 的卷积核(步长 2)做一次卷积,就得到 7×7=497\times7=49 个 dd 维 token:

patchify: [B,4,14,14]→[B,49,d]\text{patchify: } [B,4,14,14] \rightarrow [B,49,d]

Transformer 本身没有位置概念,给每个 token 加上固定的 2D sin-cos 位置编码(前一半维度编码行号,后一半编码列号)。FLUX 实际使用 RoPE(旋转位置编码),思想上同样是注入位置信息。

16.2 adaLN-Zero:条件调制

U-Net 中条件是"投影后加到特征上";DiT 系列改用更强的 adaLN-Zero(adaptive LayerNorm):

  • LayerNorm 不带可学习仿射参数(elementwise_affine=False);
  • 由条件向量(时间 tt 与条件的嵌入之和)回归出每层的调制参数——每个 block 6 组:attention 前的 shift/scale/gate 和 MLP 前的 shift/scale/gate;
  • 对 LayerNorm 输出做 modulate(h)=h⋅(1+scale)+shift\text{modulate}(h) = h\cdot(1+\text{scale}) + \text{shift},残差分支乘上 gate;
  • 回归网络零初始化:起步时所有 shift/scale/gate 都是 0,gate=0 使残差分支不起作用,整个 block 初始等价于恒等映射,训练更稳定。

16.3 双流结构

MMDiT 的"MM"(多模态)体现在:图像 token 和文本 token 各有一套完整的参数(各自的 LayerNorm、QKV 投影、输出投影、MLP),唯一的交互点是每层中间的一次联合 attention——把两路 token 拼成一个序列做一次 softmax(QK⊤)V\text{softmax}(QK^\top)V,然后再切开各走各的:

1
2
3
4
img tokens ──┐  (各自的 adaLN、QKV)      ┌──► img tokens
├──► 拼接 → 联合 attention ──┤
txt tokens ──┘ (无参数) └──► txt tokens
两套独立参数(proj、MLP 也是两套)

这样文本信息通过 attention 流入图像 token,图像 token 的梯度也能反过来塑造文本表示。本实现中"文本"仍只有 1 个类别 token,但结构与 FLUX 完全同构——换成 T5 输出的多 token 文本序列即可直接扩展。

16.4 Final layer 与 unpatchify

最后再用一次 adaLN 调制,然后线性投影回每个 patch 的像素(d→p2⋅cd \rightarrow p^2\cdot c),reshape 还原为 [B,4,14,14][B,4,14,14] 的速度预测 v^\hat v。投影层同样零初始化,让整个网络初始输出为 0(预测"不动"的速度)。

17. Mini FLUX 的 MNIST 完整实现

单文件实现:tiny AE + Flow Matching 训练 + 双流 MMDiT + Euler/CFG 采样。
环境要求:torch、torchvision。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
# mini_flux_mnist.py
# 从零实现的 Mini FLUX:
# 继承 mini_sd 的 tiny AE / 类别条件 / CFG
# 换成 Flow Matching(v 预测)+ Euler ODE 采样 + 双流 MMDiT 主干
import os
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
from torchvision.utils import save_image

# ---------------- 超参数 ----------------
BATCH_SIZE = 128
AE_EPOCHS = 2 # 自编码器预训练轮数
DIFF_EPOCHS = 10 # 扩散模型训练轮数
AE_LR, DIFF_LR = 1e-3, 3e-4
LATENT_CH = 4 # 潜空间 28x28x1 -> 14x14x4
D_MODEL = 128 # Transformer 隐藏维度
N_HEADS = 4
DEPTH = 4 # MMDiT block 数
PATCH = 2 # patch 大小(14x14 -> 49 个 token)
NUM_CLASSES = 10
NULL_CLASS = NUM_CLASSES # CFG 空条件
CFG_DROP_PROB = 0.1
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
SAVE_DIR = "output/mini_flux"
os.makedirs(SAVE_DIR, exist_ok=True)


# ---------------- 自编码器(与 mini_sd 相同)----------------
class TinyAutoencoder(nn.Module):
"""28x28x1 <-> 14x14x4,MSE 重建训练"""
def __init__(self, latent_ch=4):
super().__init__()
self.encoder = nn.Sequential(
nn.Conv2d(1, 32, 3, stride=2, padding=1), # 28 -> 14
nn.SiLU(),
nn.Conv2d(32, latent_ch, 3, padding=1),
)
self.decoder = nn.Sequential(
nn.Conv2d(latent_ch, 32, 3, padding=1),
nn.SiLU(),
nn.ConvTranspose2d(32, 1, 4, stride=2, padding=1), # 14 -> 28
nn.Tanh(),
)

def forward(self, x):
return self.decoder(self.encoder(x))


# ---------------- 时间嵌入与 2D 位置编码 ----------------
class SinusoidalTimeEmbedding(nn.Module):
def __init__(self, dim):
super().__init__()
self.dim = dim

def forward(self, t):
half = self.dim // 2
freqs = torch.exp(-math.log(10000) * torch.arange(half, device=t.device) / half)
args = t.float()[:, None] * freqs[None, :]
return torch.cat([torch.sin(args), torch.cos(args)], dim=-1)


def sincos_pos_2d(dim, h, w):
"""2D sin-cos 位置编码:前半维度编码行号,后半编码列号。返回 [1, h*w, dim]"""
def axis(dim_half, n):
omega = torch.exp(-math.log(10000) * torch.arange(dim_half // 2) / (dim_half // 2))
out = torch.arange(n)[:, None].float() * omega[None, :]
return torch.cat([out.sin(), out.cos()], dim=-1) # [n, dim_half]

rows = axis(dim // 2, h) # [h, d/2]
cols = axis(dim // 2, w) # [w, d/2]
grid = torch.cat([
rows[:, None, :].expand(h, w, dim // 2),
cols[None, :, :].expand(h, w, dim // 2),
], dim=-1)
return grid.reshape(1, h * w, dim)


# ---------------- 双流 MMDiT ----------------
def modulate(x, shift, scale):
"""adaLN 调制: h * (1 + scale) + shift"""
return x * (1 + scale) + shift


class Stream(nn.Module):
"""一条流(图像或文本)的全部带参组件"""
def __init__(self, d, mlp_ratio=4):
super().__init__()
# LayerNorm 不带仿射参数,仿射交给 adaLN
self.norm1 = nn.LayerNorm(d, elementwise_affine=False, eps=1e-6)
self.qkv = nn.Linear(d, 3 * d)
self.attn_out = nn.Linear(d, d)
self.norm2 = nn.LayerNorm(d, elementwise_affine=False, eps=1e-6)
self.mlp = nn.Sequential(
nn.Linear(d, d * mlp_ratio), nn.GELU(), nn.Linear(d * mlp_ratio, d))
# adaLN-Zero:条件向量 -> 6 组调制参数,零初始化(起步为恒等)
self.adaLN = nn.Sequential(nn.SiLU(), nn.Linear(d, 6 * d))
nn.init.zeros_(self.adaLN[-1].weight)
nn.init.zeros_(self.adaLN[-1].bias)


class MMDiTBlock(nn.Module):
"""双流 block:图像/文本两套参数,唯一交互点是一次无参的联合 attention"""
def __init__(self, d, n_heads):
super().__init__()
self.n_heads = n_heads
self.img = Stream(d)
self.txt = Stream(d)

def _qkv(self, stream, x, cond):
"""一条流的前半:adaLN 调制 + QKV 投影"""
s1, c1, g1, s2, c2, g2 = stream.adaLN(cond)[:, None, :].chunk(6, dim=-1)
h = modulate(stream.norm1(x), s1, c1)
q, k, v = stream.qkv(h).chunk(3, dim=-1)
return q, k, v, (g1, s2, c2, g2)

def _out(self, stream, x, attn, gates):
"""一条流的后半:输出投影 + 门控残差 + MLP"""
g1, s2, c2, g2 = gates
x = x + g1 * stream.attn_out(attn)
x = x + g2 * stream.mlp(modulate(stream.norm2(x), s2, c2))
return x

def forward(self, img, txt, cond):
qi, ki, vi, gi = self._qkv(self.img, img, cond)
qt, kt, vt, gt = self._qkv(self.txt, txt, cond)

# 联合 attention:拼接两路序列,softmax(QK^T)V 本身无参数
q = torch.cat([qt, qi], dim=1)
k = torch.cat([kt, ki], dim=1)
v = torch.cat([vt, vi], dim=1)
B, N, d = q.shape
hd = d // self.n_heads
q, k, v = (u.view(B, N, self.n_heads, hd).transpose(1, 2) for u in (q, k, v))
o = F.scaled_dot_product_attention(q, k, v)
o = o.transpose(1, 2).reshape(B, N, d)

n_txt = txt.shape[1]
txt = self._out(self.txt, txt, o[:, :n_txt], gt) # 切回文本流
img = self._out(self.img, img, o[:, n_txt:], gi) # 切回图像流
return img, txt


class MiniFluxDiT(nn.Module):
"""patchify -> N x MMDiTBlock -> final layer -> unpatchify,预测速度 v"""
def __init__(self, ch=4, size=14, patch=2, d=128, n_heads=4, depth=4):
super().__init__()
self.patch, self.ch = patch, ch
self.grid = size // patch # 7
self.patch_embed = nn.Conv2d(ch, d, kernel_size=patch, stride=patch)
self.register_buffer("pos", sincos_pos_2d(d, self.grid, self.grid))

self.time_mlp = nn.Sequential(
SinusoidalTimeEmbedding(d), nn.Linear(d, d), nn.SiLU(), nn.Linear(d, d))
self.label_emb = nn.Embedding(NUM_CLASSES + 1, d) # mini 版"文本编码器"

self.blocks = nn.ModuleList([MMDiTBlock(d, n_heads) for _ in range(depth)])

# final layer:adaLN 调制 + 零初始化投影回 patch
self.final_norm = nn.LayerNorm(d, elementwise_affine=False, eps=1e-6)
self.final_adaLN = nn.Sequential(nn.SiLU(), nn.Linear(d, 2 * d))
self.final_proj = nn.Linear(d, patch * patch * ch)
nn.init.zeros_(self.final_adaLN[-1].weight)
nn.init.zeros_(self.final_adaLN[-1].bias)
nn.init.zeros_(self.final_proj.weight)
nn.init.zeros_(self.final_proj.bias)

def forward(self, z, t, y):
"""z: [B,4,14,14]; t: [B] float in [0,1]; y: [B] 类别(10 = null)"""
img = self.patch_embed(z).flatten(2).transpose(1, 2) + self.pos # [B,49,d]
txt = self.label_emb(y)[:, None, :] # [B,1,d]
cond = self.time_mlp(t * 1000) + self.label_emb(y) # adaLN 条件向量

for blk in self.blocks:
img, txt = blk(img, txt, cond)

s, c = self.final_adaLN(cond)[:, None, :].chunk(2, dim=-1)
out = self.final_proj(modulate(self.final_norm(img), s, c)) # [B,49,p*p*ch]

# unpatchify: [B,49,p*p*ch] -> [B,ch,14,14]
B, p, g, ch = z.shape[0], self.patch, self.grid, self.ch
out = out.view(B, g, g, p, p, ch).permute(0, 5, 1, 3, 2, 4)
return out.reshape(B, ch, g * p, g * p) # 预测的速度 v̂


# ---------------- Flow Matching:训练 ----------------
def sample_t(batch, device):
"""logit-normal 时间采样: t = sigmoid(N(0,1)),密度集中在中间时刻"""
return torch.sigmoid(torch.randn(batch, device=device))


def get_dataloader():
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,)), # [-1, 1]
])
dataset = datasets.MNIST(root="./data", train=True, download=True, transform=transform)
return DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=True,
num_workers=2, drop_last=True)


def train():
loader = get_dataloader()
ae = TinyAutoencoder(LATENT_CH).to(DEVICE)
model = MiniFluxDiT(LATENT_CH, 14, PATCH, D_MODEL, N_HEADS, DEPTH).to(DEVICE)

# ---- 阶段一:训练并冻结自编码器(与 mini_sd 完全相同)----
opt_ae = torch.optim.AdamW(ae.parameters(), lr=AE_LR)
for epoch in range(AE_EPOCHS):
total = 0.0
for x, _ in loader:
x = x.to(DEVICE)
loss = F.mse_loss(ae(x), x)
opt_ae.zero_grad(); loss.backward(); opt_ae.step()
total += loss.item()
print(f"[AE] epoch {epoch+1}/{AE_EPOCHS} recon_loss={total/len(loader):.5f}")

ae.requires_grad_(False).eval()
with torch.no_grad():
x, _ = next(iter(loader))
latent_std = ae.encoder(x.to(DEVICE)).std().item()
print(f"latent_std = {latent_std:.3f}")

# ---- 阶段二:Flow Matching 训练 ----
opt = torch.optim.AdamW(model.parameters(), lr=DIFF_LR)
for epoch in range(DIFF_EPOCHS):
total = 0.0
for x, y in loader:
x, y = x.to(DEVICE), y.to(DEVICE)
with torch.no_grad():
z0 = ae.encoder(x) / latent_std # 像素 -> 潜空间
# CFG 条件丢弃
drop = torch.rand(y.shape[0], device=DEVICE) < CFG_DROP_PROB
y = torch.where(drop, torch.full_like(y, NULL_CLASS), y)

t = sample_t(z0.shape[0], DEVICE) # logit-normal
eps = torch.randn_like(z0)
t_bc = t.view(-1, 1, 1, 1)
zt = (1 - t_bc) * z0 + t_bc * eps # 直线插值
v_target = eps - z0 # 速度目标

v_pred = model(zt, t, y)
loss = F.mse_loss(v_pred, v_target)

opt.zero_grad(); loss.backward(); opt.step()
total += loss.item()
print(f"[FM] epoch {epoch+1}/{DIFF_EPOCHS} loss={total/len(loader):.4f}")

torch.save({"model": model.state_dict(), "ae": ae.state_dict(),
"latent_std": latent_std}, f"{SAVE_DIR}/mini_flux.pt")
sample(model, ae, latent_std, torch.arange(10, device=DEVICE),
steps=20, w=3.0, fname=f"{SAVE_DIR}/epoch_{epoch+1:03d}.png")


# ---------------- Euler ODE 采样 + CFG ----------------
@torch.no_grad()
def sample(model, ae, latent_std, y, steps=20, w=3.0, fname=None):
"""
解 ODE dz/dt = v_θ(z,t,c),从 t=1(纯噪声)积分到 t=0(数据)。
steps: Euler 步数(FLUX 式采样,20 步即可)
w: CFG 引导强度
"""
model.eval()
B = y.shape[0]
y_null = torch.full_like(y, NULL_CLASS)
z = torch.randn(B, LATENT_CH, 14, 14, device=DEVICE) # t = 1 处的纯噪声
dt = 1.0 / steps

for i in range(steps):
t = torch.full((B,), 1.0 - i * dt, device=DEVICE) # 连续时间
# CFG:条件/无条件各前向一次,外推(与 mini_sd 公式一字不差)
v_cond = model(z, t, y)
v_uncond = model(z, t, y_null)
v = v_uncond + w * (v_cond - v_uncond)
z = z - v * dt # Euler 步,逆着 t 走

x = ae.decoder(z * latent_std) # 潜空间 -> 像素
x = (x.clamp(-1, 1) + 1) / 2
if fname:
save_image(x, fname, nrow=B)
model.train()
return x


if __name__ == "__main__":
train()

运行方式:

1
python mini_flux.py

与 mini_sd 相同的训练流程:先训 AE,再训扩散模型。区别是采样只需 20 步
(mini_sd 的 DDIM 需要 50 步)即可获得相当的质量——这正是 Flow Matching
直线路径带来的加速。

训练完成后的条件生成结果(0~9 每类各一张,20 步 Euler,w=3.0w=3.0):

Mini FLUX 生成结果

18. 代码与公式的对应关系(Mini FLUX)

公式 代码位置
t=sigmoid(u), u∼N(0,1)t = \text{sigmoid}(u),\ u\sim\mathcal{N}(0,1) sample_t()
zt=(1−t)z0+tϵz_t = (1-t)z_0 + t\epsilon(直线插值) train() 中的 zt
v=ϵ−z0v = \epsilon - z_0 train() 中的 v_target
LFM=∣vθ(zt,t,c)−v∣2\mathcal{L}_{FM} = |v_\theta(z_t,t,c) - v|^2 F.mse_loss(v_pred, v_target)
patchify [B,4,14,14]→[B,49,d][B,4,14,14]\to[B,49,d] patch_embed(2×22\times2 卷积)+ flatten
modulate(h)=h(1+scale)+shift\text{modulate}(h) = h(1+\text{scale})+\text{shift} modulate()
adaLN-Zero 6 组调制参数 Stream.adaLN(输出 6d6d,零初始化)
双流联合 attention MMDiTBlock.forward() 中的 torch.cat + SDPA
Euler 步 z←z−vθ Δtz \leftarrow z - v_\theta\,\Delta t sample() 中的 z = z - v * dt
CFG v~=v∅+w(vc−v∅)\tilde v = v_\varnothing + w(v_c - v_\varnothing) sample() 中的 v_uncond + w * (v_cond - v_uncond)

三代实现横向对比

三部分实现共享同一 MNIST 实验平台,每一代只改动少量组件,便于对照学习:

组件 DDPM(第一部分) Mini SD(第二部分) Mini FLUX(第三部分)
扩散空间 像素 28×28×128\times28\times1 潜空间 14×14×414\times14\times4 潜空间 14×14×414\times14\times4
时间定义 离散 t∈{1..1000}t\in\{1..1000\} 离散 t∈{1..1000}t\in\{1..1000\} 连续 t∈[0,1]t\in[0,1],logit-normal 采样
加噪/插值 αˉtx0+1−αˉtϵ\sqrt{\bar\alpha_t}x_0+\sqrt{1-\bar\alpha_t}\epsilon 同左(x→zx\to z) (1−t)z0+tϵ(1-t)z_0+t\epsilon(直线)
预测目标 噪声 ϵ\epsilon 噪声 ϵ\epsilon 速度 v=ϵ−z0v=\epsilon-z_0
主干网络 U-Net 条件 U-Net 双流 MMDiT
条件 无 类别 embedding 加法注入 类别 embedding,双流 + adaLN
CFG 无 训练丢弃 + 双前向外推 同左(公式不变)
采样器 DDPM 1000 步 / DDIM 50 步 DDIM 50 步 + CFG Euler 20 步 + CFG
典型推理次数 1000 / 50 50 × 2(CFG 双前向) 20 × 2(CFG 双前向)

建议的对照实验:

  1. 同一初始噪声、不同实现:观察条件控制(w)与采样步数对质量的影响曲线;
  2. mini_sd 用 20 步 DDIM vs mini_flux 用 20 步 Euler:直观感受直线路径对少步采样的友好程度;
  3. 把 mini_flux 的 MMDiT 换回 CondUNet(保持 Flow Matching 不变):隔离"训练目标"与"主干网络"两个变量的贡献。

常见问题与调优方向

DDPM/DDIM(第一部分)

  1. 采样速度慢:DDPM 需要 T=1000T=1000 步逐步去噪。直接使用第 6 节的 DDIM(跳步采样,10~50 步即可),或进一步学习更先进的加速采样器(如 DPM-Solver)。
  2. 生成质量差:增大 U-Net 通道数、延长训练、使用 EMA(指数滑动平均)权重采样,都有明显提升。
  3. 损失不降:检查图像是否归一化到 [−1,1][-1,1](与噪声分布匹配),检查时间步是否作为条件正确注入网络。
  4. 条件生成:第一部分的实现是无条件生成。把类别标签 embedding 以与时间嵌入相同的方式加入网络,即可生成指定数字(这正是第二部分的起点)。

Mini SD(第二部分)

  1. 生成结果与类别不符:增大 CFG 的 w(如 3.0 → 5.0);检查条件是否真正注入(把 w 设为 0 对比,若结果无变化说明条件没起作用)。
  2. w 太大图像过饱和/伪影:这是 CFG 的已知问题,可降低 w,或对 z0_pred 做 clamp。
  3. 重建模糊:AE 容量太小或训练不足。可加深 AE、增加 AE_EPOCHS,或把 MSE 换成感知损失;真实 SD 用 KL-VAE + GAN 损失就是这个原因。
  4. 潜变量尺度异常:跳过 latent_std 归一化会导致加噪调度失配,损失难降、采样失败。更换 AE 结构后必须重新估计 latent_std。
  5. 想生成文字描述而非类别:把 label_emb 换成 CLIP/BERT 等文本编码器,把条件注入从"投影相加"换成 cross-attention,即得到真正的文生图结构(参看第三部分的双流结构)。
  6. 从 MNIST 迁移到真实图像:结构无需改变,只需要更强的 AE(下采样 4~8 倍)、更深的 U-Net(多级分辨率 + attention)和更大的数据。

Mini FLUX(第三部分)

  1. 与 mini_sd 对比实验:同一 AE、同一 CFG,只换训练目标/采样器/主干。可以观察:20 步 Euler 与 50 步 DDIM 的质量差异;MMDiT 与 U-Net 的收敛速度差异。
  2. 采样步数:Flow Matching 的直线路径让 5~10 步也能得到可辨认结果,可以试 steps=5 对比 steps=50,直观理解"路径越直、离散误差越小"。
  3. 训练不稳定:adaLN-Zero 和 final_proj 的零初始化是稳定性的关键,去掉后 loss 初期会剧烈震荡——可以做个消融验证。
  4. t 的采样分布:把 logit-normal 换成均匀分布,中间噪声水平训练不足,采样中段质量下降明显(又一个可做的消融)。
  5. 扩展到真实文生图:类别 embedding 换成 T5 编码的多 token 文本序列(文本流长度从 1 变为数百)、sin-cos 位置编码换成 RoPE、加深/加宽 MMDiT——即得到 FLUX 的完整结构。
  6. 更快的采样:Euler 是一阶方法,换成 Heun/中点法(二阶)可以在更少步数下保持质量,这也是各开源 FLUX pipeline 的常见优化。

参考文献

  • Ho et al., Denoising Diffusion Probabilistic Models (DDPM), NeurIPS 2020
  • Song et al., Denoising Diffusion Implicit Models (DDIM), ICLR 2021
  • Rombach et al., High-Resolution Image Synthesis with Latent Diffusion Models (LDM / Stable Diffusion), CVPR 2022
  • Ho & Salimans, Classifier-Free Diffusion Guidance, NeurIPS 2021 Workshop
  • Radford et al., Learning Transferable Visual Models From Natural Language Supervision (CLIP), ICML 2021
  • Esser et al., Scaling Rectified Flow Transformers for High-Resolution Image Synthesis (SD3 / MMDiT), ICML 2024
  • Liu et al., Flow Straight and Fast: Learning to Generate and Transfer Data with Rectified Flow, ICLR 2023
  • Lipman et al., Flow Matching for Generative Modeling, ICLR 2023
  • Peebles & Xie, Scalable Diffusion Models with Transformers (DiT / adaLN-Zero), ICCV 2023
  • Black Forest Labs, FLUX.1, 2024
  • Luo, Understanding Diffusion Models: A Unified Perspective, 2022

扩散模型MNIST实战
https://huan-yin.github.io/2026/09/06/扩散模型MNIST实战/
作者
李相越
发布于
2026年9月6日
许可协议