📑 本页目录(点开跳转)
03 · 训练目标为什么是「预测噪声」
⏱ 30 分钟 | ⭐ 一串 KL 塌成一行均方误差,而标准答案是你自己造的
🎯 一句话
训练扩散模型的损失就是一行均方误差,而它的「标准答案」正是你自己刚加进去的那份噪声。
上一章停在「估噪声和估原图数学上是同一件事」。这一章补完两问:那行均方误差从哪来,以及既然等价、为什么一律选噪声。
🏷️ 一、先破一个误解:标签是你自己造的
「扩散模型」听着像某种无监督黑魔法。它不是。训练一步只有五个动作:
- 取一张真图 x₀,随机抽一个时刻 t(1…T 均匀)
- ⭐ 随机采一份噪声 ε ~ N(0, I) —— 这一份就是标签
- 闭式算出 $x_t = \sqrt{\bar\alpha_t}x_0 + \sqrt{1-\bar\alpha_t}\,\varepsilon$
- 把 (x_t, t) 喂进网络,拿到猜测 ε_θ(x_t, t)
- 损失 ‖ε − ε_θ‖²,反传
⭐ 第 2 步采的和第 5 步比的是同一个 ε —— 不是待估的未知量,是你几行之前亲手
randn出来的那个张量。所以这从头到尾是监督回归:不用蒙特卡洛,不用跑完整条链。
⭐ 也因此训练数据几乎无限:同一张图配不同的 t 和 ε,就是一条全新样本。
📐 二、那行均方误差是从哪来的
它是从变分下界化简掉下来的。骨架三句:
- x₁…x_T 全是隐变量,log p(x₀) 要对中间状态积分,算不了;改为最大化它的下界 ELBO。
- 展开后下界是 T 项 KL 之和,每项都是「真实的那一步」和「模型走的那一步」的距离。
- ⭐ 关键一刀:两边都是高斯,方差写死、只有均值要学。同方差高斯的 KL 等于均值差的平方除以 2σ²,一串 KL 当场塌成一串平方差。
再把上一章 μ̃_t 的形状代进去,x_t 两边抵消,只剩噪声那一项:
$$L_{t-1} = \frac{\beta_t^2}{2\sigma_t^2\alpha_t(1-\bar\alpha_t)}\,\big\|\varepsilon - \varepsilon_\theta(x_t,t)\big\|^2$$
⭐ DDPM 原文再补一刀:把系数直接扔掉,只留 L_simple = E‖ε − ε_θ‖²。它在 t 小时特别大,扔掉等于给 t 大的样本更高权重(那段决定整体构图),实测出图反而更好。
📐 完整推导(想看再点)
按马尔可夫链把 Jensen 给出的证据下界拆开、用贝叶斯把每步换成以 x₀ 为条件的形式,得
$$L_{\text{VLB}} = D_{KL}\big(q(x_T|x_0)\|p(x_T)\big) + \sum_{t\ge 2} D_{KL}\big(q(x_{t-1}|x_t,x_0)\|p_\theta(x_{t-1}|x_t)\big) - \log p_\theta(x_0|x_1)$$
首项无可学参数(扔),末项是收尾重建,中间 T−1 项是全部战场。
第一刀:q(x_{t−1} | x_t, x₀) 是高斯,模型那边方差也定成常数 σ_t²,这项 KL 就是 ‖μ̃_t − μ_θ‖² ÷ 2σ_t²。
第二刀:把 x₀ 用 x_t 和 ε 表示代回上一章那个 μ̃_t,化简后它变成 (1/√α_t)·(x_t − β_t/√(1−ᾱ_t)·ε);让模型照同一形状写 μ_θ。⭐ 相减时 x_t 抵消,只剩 ε − ε_θ 乘一个常数。
⚠️ 「预测噪声」正是在这一刀选定的:吐 x̂₀ 或 μ_θ 也走得完,只是平方差里装的东西不同。
⚖️ 三、三种目标可以互换,为什么偏偏选噪声
给我任意一个,都能算出另外两个:
$$\hat x_0 = \frac{x_t - \sqrt{1-\bar\alpha_t}\,\varepsilon_\theta}{\sqrt{\bar\alpha_t}}, \qquad v = \sqrt{\bar\alpha_t}\,\varepsilon - \sqrt{1-\bar\alpha_t}\,x_0$$
既然能互换,选哪个应该无所谓才对。但损失长在哪个量上,梯度就长在哪个量上:
# 同一份预测误差,换成两种训练目标之后,损失量级差多少
import torch
torch.manual_seed(0)
abar = torch.cumprod(1 - torch.linspace(1e-4, 0.02, 1000), 0) # ᾱ_t,linear 调度
x0 = torch.randn(200000) # 干净数据(方差 1)
for t in [10, 500, 990]:
a = abar[t - 1]
eps = torch.randn_like(x0)
xt = a.sqrt() * x0 + (1 - a).sqrt() * eps # 第 2 章那条闭式加噪
eh = eps + 0.1 * torch.randn_like(eps) # ⭐ 网络水平固定:ε 上恒差 0.1
xh = (xt - (1 - a).sqrt() * eh) / a.sqrt() # 同一个预测换算成 x̂₀
print(f"t={t:4d} k={(1-a).sqrt()/a.sqrt():6.2f} |"
f" 摆烂猜x0={((xt-x0)**2).mean():6.4f} 摆烂猜eps={(eps**2).mean():6.4f} |"
f" x0-MSE={((xh-x0)**2).mean():8.2e} eps-MSE={((eh-eps)**2).mean():6.4f}")
真实输出(摆烂 = 什么都不学的网络:猜 x₀ 就原样吐回输入,猜 ε 就一律说「没加噪」):
t= 10 k= 0.04 | 摆烂猜x0=0.0019 摆烂猜eps=1.0024 | x0-MSE=1.90e-05 eps-MSE=0.0100
t= 500 k= 3.42 | 摆烂猜x0=1.4401 摆烂猜eps=0.9965 | x0-MSE=1.17e-01 eps-MSE=0.0100
t= 990 k=142.35 | 摆烂猜x0=1.9877 摆烂猜eps=0.9995 | x0-MSE=2.02e+02 eps-MSE=0.0100
- ① 预测 x₀ 的难度随 t 剧烈摆动,预测 ε 不会。 t=10 时 x_t 和原图几乎一样,摆烂网络只错 0.0019(⭐ 损失小 = 梯度小 = 几乎不给学习信号);t=990 反过来要错 1.9877。而摆烂猜 ε 两头是 1.0024 和 0.9995,一分钱没占到。
- ② 后两列是同一个网络(ε 上恒差 0.1):换成 x₀ 口径,损失从 1.90e-05 跨到 202,七个数量级;ε 口径恒定 0.0100。换算比 k = √(1−ᾱ_t) ÷ √ᾱ_t 在 t=990 时是 142.35。
⭐⭐ 所以选目标不是数学问题,是梯度量级问题。 对 x₀ 做均方误差,梯度会被少数 t 大的样本吃光;预测 ε 把这条曲线拉平了。
| 参数化 | 强在哪 | 弱在哪 |
|---|---|---|
| ε(默认) | 各个 t 的损失量级均衡;DDPM / DDIM 默认 | t 接近 0 时反推 x̂₀ 要除以 √ᾱ_t,误差放大 142 倍 |
| x₀ | t 大时稳,不用除小数 | t 小时目标几乎等于输入,学不到东西 |
| v:按噪声水平把 ε 和 x₀ 旋转混合 | ⭐ 两头都不塌;少步蒸馏、zero terminal SNR 基本必用 | 多一层换算,读代码容易绕晕 |
⏰ 四、时间步 t 怎么告诉网络
同一份权重要在 1000 个噪声水平上干活,而 t=10 该做的事(微调纹理)和 t=990 该做的事(从一团雾里定构图)完全不同。所以 t 必须进网络,而且要送到每一层。
⚠️ 但不能直接喂那个整数:0~999 的标量信息量太小,尺度也和特征图上的激活值格格不入。
⭐ 做法和 Transformer 给 token 编位置是同一个套路:一组频率各异的正弦余弦把整数摊成高维向量,再过两层小 MLP,加到每个残差块的特征上。
$$\text{emb}(t)_{2i} = \sin(t\,\omega_i), \quad \text{emb}(t)_{2i+1} = \cos(t\,\omega_i), \quad \omega_i = 10000^{-2i/d}$$
人话:低频那几维几百步才转一圈(管「在链条哪一段」),高频那几维一步一变(管「精确第几步」),合起来每个 t 拿到一串独一无二的条纹指纹。好处有二:相邻的 t 得到相近的向量,网络能插值,不会把 1000 个时刻当成 1000 个互不相干的类别;而且不用学,换 T 也照样算。
⚠️ 代码里它一般叫 timestep_embedding 或 time_mlp;连续时间的实现喂的不是 t 而是噪声水平 σ 或 log-SNR。
🔗 这一章连到哪里
| 相关的地方 | 为什么 |
|---|---|
| 数学原理 01b · KL 散度 | 第二节那串 KL 的地基。⭐ 还补了本章没展开的一条:ELBO 里那项是反向 KL,「允许 Q 只挑一个峰蹲着」 |
| 数学原理 13 · EM 与高斯混合 | ELBO 的正式出处:E 步把下界顶到贴着似然,M 步抬高下界。扩散是同一招,隐变量换成「中间那 T 张带噪图」 |
| ML 基础 09 · 优化器与学习率 | 「梯度被少数样本吃光」是这一层的事。那一章讲 Adam 怎么给每个参数各自的步长来救场 |
| 全景导论 专题 C · RoPE:位置感 | 正弦嵌入的另一个用法:「不同频率的正弦叠在一起,每个位置拿到一串独一无二的条纹指纹」——同一个零件 |
✅ 检查点
- 损失里那个 ε 是从哪来的?为什么说这是监督学习?
- ELBO 展开成一串 KL 之后,是哪个性质让它「塌成」平方差的?
- L_simple 比完整的变分下界少了什么?后果是什么?
- t=10 时,一个「原样吐回输入」的网络在 x₀ 目标上错多少?在 ε 目标上呢?说明什么?
- 同一个网络(ε 上恒差 0.1),换成 x₀ 口径后损失跨了多大范围?
- 为什么时间步不能直接喂一个整数?
👀 答案
- 是你自己
randn出来的那份噪声 —— 第 2 步采出、第 3 步拿它造 x_t、第 5 步拿它当标签,同一个张量。标签已知,所以是监督回归。 - 两边都是高斯,方差写死、只学均值。 同方差高斯的 KL 等于均值差的平方除以 2σ²;再把 μ_θ 写成和 μ̃_t 同一形状,相减时 x_t 抵消,只剩 ε − ε_θ。
- 少了系数 β_t²/(2σ_t²α_t(1−ᾱ_t))。它在 t 小时特别大,扔掉等于给 t 大的样本更高权重,而那段决定整体构图。
- x₀ 上 0.0019(几乎白送),ε 上 1.0024(一分钱没占到):t 小时预测 x₀ 太容易,损失小 = 梯度小 = 几乎不给学习信号。
- 从 1.90e-05 跨到 202,七个数量级;ε 口径恒定 0.0100。换算比 k 在 t=990 时是 142.35。
- 0~999 的标量信息量太小,尺度也和激活值对不上。正弦嵌入摊成高维向量:低频维管「在链条哪一段」、高频维管「精确第几步」,相邻 t 向量相近所以能插值。
🛑 可以停在这里
读到这里,你已经能看懂任意一份扩散训练脚本的核心几行:t = torch.randint(0, T, (B,)) 是每张图各抽各的时刻,noise = torch.randn_like(x0) 就是标签,F.mse_loss(model(x_t, t), noise) 就是 L_simple,而 prediction_type: v_prediction 只是换了参数化、不是换了模型。
什么时候回来:改了预测目标却忘了改采样端的换算(第三节);loss 好看但出图糊(第二节);想上 zero terminal SNR 却发现 ε 参数化在终点除零(第三节的表)。
⚡ 走神救援
先记住这几件事
- 训练样本由干净数据、随机时刻和已知噪声构造,目标可以直接计算。
- 预测噪声、预测干净样本和 v 参数化之间可以换算,但损失的时间权重并不因此相同。
- 时间条件必须传给去噪网络;先核对采样出来的噪声与监督目标是同一份。
下一节 👉 04-DDIM与采样加速.md