🏠 总目录📚 本教程 03 · 训练目标为什么是预测噪声 ← →
📑 本页目录(点开跳转)

03 · 训练目标为什么是「预测噪声」

⏱ 30 分钟 | ⭐ 一串 KL 塌成一行均方误差,而标准答案是你自己造的


🎯 一句话

训练扩散模型的损失就是一行均方误差,而它的「标准答案」正是你自己刚加进去的那份噪声。

上一章停在「估噪声和估原图数学上是同一件事」。这一章补完两问:那行均方误差从哪来,以及既然等价、为什么一律选噪声。


🏷️ 一、先破一个误解:标签是你自己造的

「扩散模型」听着像某种无监督黑魔法。它不是。训练一步只有五个动作:

  1. 取一张真图 x₀,随机抽一个时刻 t(1…T 均匀)
  2. ⭐ 随机采一份噪声 ε ~ N(0, I) —— 这一份就是标签
  3. 闭式算出 $x_t = \sqrt{\bar\alpha_t}x_0 + \sqrt{1-\bar\alpha_t}\,\varepsilon$
  4. 把 (x_t, t) 喂进网络,拿到猜测 ε_θ(x_t, t)
  5. 损失 ‖ε − ε_θ‖²,反传

⭐ 第 2 步采的和第 5 步比的是同一个 ε —— 不是待估的未知量,是你几行之前亲手 randn 出来的那个张量。所以这从头到尾是监督回归:不用蒙特卡洛,不用跑完整条链。

⭐ 也因此训练数据几乎无限:同一张图配不同的 t 和 ε,就是一条全新样本。


📐 二、那行均方误差是从哪来的

它是从变分下界化简掉下来的。骨架三句:

再把上一章 μ̃_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 的损失量级均衡;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 也照样算。

t 不是一个数字,要变成一串能被每一层读到的向量t = 990一个整数不同频率的正弦低频管「大概在哪一段」,高频管「精确是第几步」时间嵌入一串浮点数每一层都注入⭐ 和位置编码是同一个套路:把「第几个」变成一串不同频率的正弦⚠️ 只在输入端塞一次不够 —— t 要影响每一层的行为
看中间那三条**频率不同**的正弦:低频区分「大致在哪一段」,高频区分「精确是第几步」。⭐ 和 Transformer 的位置编码是同一个套路,只是这里编的是时间而不是位置。

⚠️ 代码里它一般叫 timestep_embedding 或 time_mlp;连续时间的实现喂的不是 t 而是噪声水平 σ 或 log-SNR。


🔗 这一章连到哪里

相关的地方 为什么
数学原理 01b · KL 散度 第二节那串 KL 的地基。⭐ 还补了本章没展开的一条:ELBO 里那项是反向 KL,「允许 Q 只挑一个峰蹲着」
数学原理 13 · EM 与高斯混合 ELBO 的正式出处:E 步把下界顶到贴着似然,M 步抬高下界。扩散是同一招,隐变量换成「中间那 T 张带噪图」
ML 基础 09 · 优化器与学习率 「梯度被少数样本吃光」是这一层的事。那一章讲 Adam 怎么给每个参数各自的步长来救场
全景导论 专题 C · RoPE:位置感 正弦嵌入的另一个用法:「不同频率的正弦叠在一起,每个位置拿到一串独一无二的条纹指纹」——同一个零件

✅ 检查点

  1. 损失里那个 ε 是从哪来的?为什么说这是监督学习?
  2. ELBO 展开成一串 KL 之后,是哪个性质让它「塌成」平方差的?
  3. L_simple 比完整的变分下界少了什么?后果是什么?
  4. t=10 时,一个「原样吐回输入」的网络在 x₀ 目标上错多少?在 ε 目标上呢?说明什么?
  5. 同一个网络(ε 上恒差 0.1),换成 x₀ 口径后损失跨了多大范围?
  6. 为什么时间步不能直接喂一个整数?
👀 答案
  1. 是你自己 randn 出来的那份噪声 —— 第 2 步采出、第 3 步拿它造 x_t、第 5 步拿它当标签,同一个张量。标签已知,所以是监督回归。
  2. 两边都是高斯,方差写死、只学均值。 同方差高斯的 KL 等于均值差的平方除以 2σ²;再把 μ_θ 写成和 μ̃_t 同一形状,相减时 x_t 抵消,只剩 ε − ε_θ。
  3. 少了系数 β_t²/(2σ_t²α_t(1−ᾱ_t))。它在 t 小时特别大,扔掉等于给 t 大的样本更高权重,而那段决定整体构图。
  4. x₀ 上 0.0019(几乎白送),ε 上 1.0024(一分钱没占到):t 小时预测 x₀ 太容易,损失小 = 梯度小 = 几乎不给学习信号。
  5. 从 1.90e-05 跨到 202,七个数量级;ε 口径恒定 0.0100。换算比 k 在 t=990 时是 142.35。
  6. 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 却发现 ε 参数化在终点除零(第三节的表)。

⚡ 走神救援

先记住这几件事

下一节 👉 04-DDIM与采样加速.md

打卡记录保存在你的浏览器里,首页能看到总进度