AIGC 基本功|VAE 结构与训练目标-VAE

人工智能炼丹君
2026-09-26 / 0 评论 / 3 阅读 / 正在检测是否收录...

VAE 结构、KL 权重,与那个神秘的 0.18215

所属方向:表征与压缩 | 难度:进阶 | 前置知识:变分下界与重参数化(本篇是它的直接后继,ELBO 与重参数化在那里推过,这里只用结论)
关键词:VAE、编码器、解码器、潜空间、KL 权重、后验坍缩、scaling factor


01. 为什么需要它

先看一个可复现的尺度错误:预训练扩散模型要求 VAE latent 乘以 vae.config.scaling_factor,却把原始 latent 直接送进去。张量形状仍正确,但它已经偏离训练分布。下面用一维高斯去噪器量化这个误差;这些数字来自教学模型,不是 Stable Diffusion 图像质量实验,不能据此断言真实图像一定发灰或过曝。

这不是玄学,可以精确算出来。设潜空间真实标准差是 $\sigma$,扩散模型的噪声表却是按「数据方差为 1」标定的。在完全干净那一端($\bar{\alpha} \to 1$)两边都对;越往噪声端走,误差越大。用高斯 MMSE 估计可以算出,模型的去噪幅度只有正确值的

$$r(\bar{\alpha}) = \bar{\alpha} + \frac{1 - \bar{\alpha}}{\sigma^{2}}$$

倍。$\sigma = 5.49$(这是由 SD1.x 常用配置系数反推的尺度,第 03.4 节会讲为什么)时,$\bar{\alpha} = 0.5$ 处 $r = 0.5166$——幅度只剩一半。实测的均方误差从本该有的 0.9682 涨到 7.8154,恶化 8.07 倍。这是高斯教学模型的估计误差,不是图像实测。

反着错也一样疼。如果你的 VAE 是规规矩矩训的(潜空间方差约等于 1),却照抄 Stable Diffusion 的 0.18215,那么 $\sigma$ 变成 0.18215,同一个 $\bar{\alpha}=0.5$ 处 $r = 15.5699$——幅度被放大 15.6 倍,教学去噪器的后验均值幅度被高估。这个常数抄错方向,比抄错符号更常见。

还有第三种错法,比前两种隐蔽得多。LDM 论文(arXiv:2112.10752)附录 D.1 里有一段原话,把它写得很清楚:

the signal-to-noise ratio induced by the variance of the latent space (i.e. $\text{Var}(z)/\sigma_t^{2}$) significantly affects the results for convolutional sampling ... when training a LDM directly in the latent space of a KL-regularized model, this ratio is very high, such that the model allocates a lot of semantic detail early on in the reverse denoising process ... Note that the VQ-regularized space has a variance close to 1, such that it does not have to be rescaled.

这段观察针对 LDM 论文中具体的 KL / VQ 自编码器:KL latent 的高方差改变了给定噪声调度下的信噪比,影响高分辨率卷积采样;论文所用 VQ latent 的方差接近 1。它不表示所有 VQ 码本天然归一化,也不表示尺度错误一定对应某一种视觉伪影。

最后是 KL 权重本身的坑。在我的最小实验里,把 KL 权重 $\beta$ 从 0.1 调到 1,重建 MSE 从 0.2119 跳到 1.0000——1.0000 就是「什么都不学、直接输出均值」的分数(数据逐维方差已被归一化成 1)。潜变量的 8 个维度全部死掉。这个现象叫后验坍缩。

这一篇就把这三件事串起来:VAE 的结构决定了潜空间里有什么,KL 权重决定了还剩下什么,而剩下东西的尺度就是那个 0.18215。

02. 最小可用理解

三句话:

  1. 结构:编码器把 $x$ 压成 $2d$ 个数($d$ 个均值 $\mu$、$d$ 个对数方差 $\log\sigma^{2}$),重参数化采出一个 $z$,解码器从 $z$ 重建 $x$。训练目标是 ELBO——本质上是「重建质量」减「每个潜变量维度花掉的 KL 预算」。
  2. KL 项既是正则,也可以看作信息预算。一维潜变量要花掉多少 KL,就必须换回足够的重建收益,否则最优解就是关掉这一维。实测里 $\beta = 10^{-3}$ 时模型正好活 4 个维度(合成数据的真实因子数就是 4),$\beta$ 收到 0.3 只剩 3 个,$\beta = 1$ 一个不剩。维度是一个个死的,不是一起死。
  3. 潜空间的尺度是 KL 权重留下的痕迹。$\beta$ 大到 KL 有效时,它把潜空间边际标准差钉在 1 附近(实测 1.0067);$\beta$ 小到 KL 失效,尺度就失去约束,随训练动力学漂走(实测漂到 2.7174)。Stable Diffusion 那套 VAE 的 KL 权重是 $10^{-6}$(LDM 论文原话:we either weight the KL term by a factor $\sim 10^{-6}$),其配置对应的原始 latent 标准差约为 5.49,不能仅凭 KL 系数推断具体漂移过程——而 $1/5.49 = 0.18215$。

这张图要看什么:四张子图连起来读。(a) 重建 MSE 在 $\beta$ 超过 0.1 后逐渐上升到 1.0 那条虚线,说明模型彻底放弃潜变量;(b) 总 KL 同步归零,先验和近似后验重合;(c) 潜空间边际标准差:蓝线(无权重衰减)在 $\beta \ge 10^{-3}$ 后紧紧贴着 1,$\beta$ 一小就抬头,红线(有权重衰减)在同一条路上走得更快更远,绿虚线是 Stable Diffusion 的 5.49;(d) 还活着的维度数从 8 一路掉到 0,中间在 4 这个地方有个明显的台阶——那是合成数据的真实因子数。

03. 数学推导

3.1 一次编码到底出了什么

设数据 $x \in \mathbb{R}^{D}$,潜变量 $z \in \mathbb{R}^{d}$。VAE 的编码器不输出一个 $z$,它输出一个分布:

$$q_{\phi}(z \mid x) = \mathcal{N}\!\left(z;\ \mu_{\phi}(x),\ \text{diag}\left(\sigma_{\phi}^{2}(x)\right)\right)$$

网络实际算出来的是 $\mu$ 和对数方差 $\text{logvar}$ 两组数,各 $d$ 个。为什么是 logvar 而不是方差:网络输出可以取任意实数,但方差必须为正。把网络的输出过一层 $\exp$ 就得到恒正的方差,同时把它放进对数域还有一个好处——数值范围。实验中 $\beta = 10^{-6}$ 时后验标准差会掉到 0.009 量级,方差是 $8\times 10^{-5}$;如果网络直接回归方差,这个量级上梯度会和重建项的梯度混在一起互相淹没,换到对数域后它就只是一个普通的负实数。

从 $q_{\phi}(z|x)$ 里采一个 $z$ 直接用会断掉梯度(采样操作不可导)。重参数化把它挪出去:

$$z = \mu + \sigma \odot \epsilon, \qquad \epsilon \sim \mathcal{N}(0, I)$$

$\odot$ 是逐元素乘。随机性全部塞进 $\epsilon$ 里,$\mu$ 和 $\sigma$ 变成普通可导函数——这是整篇 VAE 能训起来的前提。

解码器给的是似然,取对角高斯:

$$p_{\theta}(x \mid z) = \mathcal{N}\!\left(x;\ \text{dec}_{\theta}(z),\ \sigma_{\text{dec}}^{2} I\right)$$

取对数:

$$\log p_{\theta}(x \mid z) = -\frac{\lVert x - \text{dec}_{\theta}(z) \rVert^{2}}{2\sigma_{\text{dec}}^{2}} - \frac{D}{2}\log\left(2\pi\sigma_{\text{dec}}^{2}\right)$$

第二项和 $z$ 无关,是常数。所以重建项就是平方误差之和除以 $2\sigma_{\text{dec}}^{2}$——「重建用 MSE」不是拍脑袋定的,它是高斯似然的必然结果,而 $\sigma_{\text{dec}}$ 就是重建项的隐式权重。日志里那种 recon + beta * kl 的写法,还要明确 reduction:若 recon 是逐元素均值,则标准负 ELBO 按同尺度缩放后的 KL 系数为 $2\sigma_{\text{dec}}^{2}/D$。

把两件事拼起来,ELBO 说 $\log p(x) \ge \mathbb{E}_{q}\left[\log p_{\theta}(x|z)\right] - \text{KL}\!\left(q_{\phi}(z|x)\,\Vert\,p(z)\right)$,先验取标准正态 $p(z) = \mathcal{N}(0, I)$。我们要最小化的就是负 ELBO:

$$J = \frac{\mathbb{E}_{q}\left[\lVert x - \text{dec}_{\theta}(z) \rVert^{2}\right]}{2\sigma_{\text{dec}}^{2}} + \beta \sum_{j=1}^{d} \text{KL}_{j}, \qquad \text{KL}_{j} = \text{KL}\!\left(\mathcal{N}(\mu_{j}, \sigma_{j}^{2})\,\Vert\,\mathcal{N}(0,1)\right)$$

$\beta$ 是我们插进去的旋钮:$\beta = 1$ 是标准 ELBO,$\beta > 1$ 就是 $\beta$-VAE 路线,调大换解耦表征。注意 KL 是对维度求和而不是求平均,这一点第 05 节还会回来算账。

3.2 KL 的闭式解,逐项推

两个对角高斯之间的 KL 有闭式解。按定义 $\text{KL}(q \Vert p) = \mathbb{E}_{q}[\log q] - \mathbb{E}_{q}[\log p]$ 分头算。因为逐维独立,下面只看第 $j$ 维。

先算 $\mathbb{E}_{q}[\log q]$。$q_{j} = \mathcal{N}(\mu_{j}, \sigma_{j}^{2})$,所以

$$\log q_{j}(z_{j}) = -\frac{1}{2}\log(2\pi\sigma_{j}^{2}) - \frac{(z_{j} - \mu_{j})^{2}}{2\sigma_{j}^{2}}$$

对 $q_{j}$ 取期望时,右边第二项的期望是 $\sigma_{j}^{2} / (2\sigma_{j}^{2}) = 1/2$,于是

$$\mathbb{E}_{q_{j}}\left[\log q_{j}\right] = -\frac{1}{2}\left(1 + \log 2\pi + \log \sigma_{j}^{2}\right)$$

只有三项,和 $\mu_{j}$ 无关——这一点值得停一下:在 $\mu_{j}$ 上平移一个高斯分布,它的熵不变。

再算 $\mathbb{E}_{q}[\log p]$。先验 $p_{j} = \mathcal{N}(0,1)$,同样展开

$$\mathbb{E}_{q_{j}}\left[\log p_{j}\right] = -\frac{1}{2}\log 2\pi - \frac{\mathbb{E}_{q_{j}}\left[z_{j}^{2}\right]}{2}$$

这里用到 $z_{j} = \mu_{j} + \sigma_{j}\epsilon$,所以 $\mathbb{E}_{q_{j}}[z_{j}^{2}] = \mu_{j}^{2} + \sigma_{j}^{2}$——这就是 $z$ 的二阶矩,它才是把 $\mu$ 拉进公式的那一项。

两者相减,$-\frac{1}{2}\log 2\pi$ 正好抵消:

$$\text{KL}_{j} = \frac{1}{2}\left(\mu_{j}^{2} + \sigma_{j}^{2} - \log \sigma_{j}^{2} - 1\right)$$

对 $j$ 求和即得总 KL。这个式子里每一块都有明确的物理含义,逐项读:

  • $\mu_{j}^{2}$ 是均值偏离先验中心的成本。 不同输入的均值发生变化可以传递信息,但这一项不是互信息本身。即使所有输入的均值都是 0,只要方差仍依赖输入,潜变量也可能携带信息;只有整个条件分布都与输入无关时,这一维才不传信息。逐维 KL 同时包含信息代价与聚合后验偏离先验的代价,不能把 8~12 nats 直接叫作有效信息量。

3.3 权重到底在权衡什么:一维线性 VAE 的闭式解

上面的直觉可以算到精确解。把模型简化到最狠:一维数据 $x \sim \mathcal{N}(0, v)$,线性编码器 $q(z|x) = \mathcal{N}(mx, s^{2})$($m$ 是缩放系数,$s^{2}$ 是固定的后验方差),线性解码器 $p(x|z) = \mathcal{N}(wz, \sigma_{\text{dec}}^{2})$。目标函数展开成

$$J(w, m, s) = \frac{v(1 - wm)^{2} + w^{2} s^{2}}{2\sigma_{\text{dec}}^{2}} + \frac{\beta}{2}\left(m^{2} v + s^{2} - \log s^{2} - 1\right)$$

第一项里的 $v(1-wm)^{2}$ 是「编码-解码这条路的增益偏离 1 有多远」,$w^{2}s^{2}$ 是「采样噪声被放大 $w$ 倍后落在输出上的方差」;第二项就是上一节的 KL。三个未知数各求一次偏导并置零:

$$\frac{\partial J}{\partial w} = \frac{-v m(1 - wm) + w s^{2}}{\sigma_{\text{dec}}^{2}} = 0, \qquad \frac{\partial J}{\partial m} = \frac{-v w(1 - wm)}{\sigma_{\text{dec}}^{2}} + \beta m v = 0, \qquad \frac{\partial J}{\partial s} = \frac{w^{2} s}{\sigma_{\text{dec}}^{2}} + \beta\left(s - \frac{1}{s}\right) = 0$$

记路增益 $u := wm$(它就是「潜变量被真正使用的程度」),把 $\partial_m$ 的式子两边乘 $w$ 换成 $u$,可以整理出 $w^{2}(1-u) = \beta \sigma_{\text{dec}}^{2} u$;再从 $\partial_w$ 得到 $w^{2} s^{2} = v u (1-u)$;从 $\partial_s$ 得到 $s^{2}\left(\beta + 2c w^{2}\right) = \beta$(这里 $2c = 1/\sigma_{\text{dec}}^{2}$)。三式联立,令

$$\lambda := \frac{\beta \sigma_{\text{dec}}^{2}}{v}$$

解出来是一组非常干净的东西:

$$u = 1 - \lambda, \qquad s^{2} = \lambda, \qquad w^{2} = v(1 - \lambda)$$

代回验一遍:$\partial_m$ 要求 $w(1-u)/\sigma_{\text{dec}}^{2}=\beta m$。乘以 $w$,代入 $w^2=v(1-\lambda)$ 和 $1-u=\lambda$,得到 $v(1-\lambda)\lambda/\sigma_{\text{dec}}^{2}=\beta(1-\lambda)$。在未坍缩区间约去 $1-\lambda$,正好得到 $\lambda=\beta\sigma_{\text{dec}}^{2}/v$。第 04.2 节的数值优化与此吻合。

$\lambda$ 的读法:它是 KL 权重乘观测噪声方差,再除以数据方差的无量纲比值。这个比例决定一切:

  • $u = 1-\lambda$ 随 $\lambda$ 线性下降。$\lambda \to 0$ 时 $u \to 1$、$s^{2} \to 0$,退化成确定性自编码器——潜变量满负荷工作,采样噪声归零。
  • $s^{2}=\lambda$ 随权重变化;$\beta=1$ 时等于这个线性高斯模型的真实后验方差,其他权重一般对应不同的变分目标。
  • 当 $\lambda>1$ 时上述分支要求 $w^2<0$,不存在实数解;$\lambda=1$ 时它连续接到坍缩点,最优解退化成坍缩点 $(w, m, s) = (0, 0, 1)$。*坍缩阈值是 $\beta^{} = v / \sigma_{\text{dec}}^{2}$**:数据方差越大、或者解码器的观测噪声越小,需要的 $\beta$ 就越大。这是线性高斯模型的阈值,不是神经 VAE 的通用保证。
  • 附带一个值得记住的事实:$(0,0,1)$ 在任意 $\beta$ 下都是驻点(数值验证梯度恒等于 0)。它只是当 $\lambda < 1$ 时不是最小值。所以「后验坍缩」不是数值 bug、不是训练不充分——在本线性模型的阈值以上,它是该目标的最优解;一般神经模型也可能因局部最优和优化动力学而坍缩。

3.4 潜空间的统计性质,和那个 0.18215

现在把「编码器」反过来看会得到什么分布。训练完之后,把所有 $x$ 编码一遍,潜变量的边际分布是

$$q(z) = \int q_{\phi}(z \mid x)\, p(x)\, \mathrm{d}x$$

这是个混合分布。对每一维分别算方差,用全方差公式:

$$\text{Var}(z_{j}) = \underbrace{\text{Var}_{p_{\text{data}}}\!\left[\mu_{j}(x)\right]}_{\text{信息}} + \underbrace{\mathbb{E}_{p_{\text{data}}}\!\left[\sigma_{j}^{2}(x)\right]}_{\text{采样噪声}}$$

左边是 scaling_factor 要归一化的东西,右边两项来源完全不同:前一项是「不同样本被编码到不同位置」,后一项是「每个样本自己抖多少」。下表先对每维方差取平均再开方,不能先平均标准差再平方,也不能把两项标准差直接相加。方差占比不是互信息。实测拆账($\beta$ 从小到大):

$\beta$ 均值变化标准差 RMS 后验噪声标准差 RMS 边际标准差 RMS 总方差中均值变化占比
1e-06 2.060 0.009 2.060 100%
1e-05 1.411 0.014 1.411 100%
0.0001 1.004 0.355 1.065 89%
0.001 0.726 0.697 1.007 52%
0.01 0.705 0.717 1.005 49%
0.1 0.620 0.779 0.995 39%
0.3 0.407 0.914 1.000 17%
1 0.003 1.000 1.000 0%

$\beta = 10^{-3}$ 附近有个交叉点:潜空间方差里「信息」和「噪声」各占一半。往左走,$\text{std}(z)$ 几乎全部来自信息;往右走,几乎全部来自采样噪声——$\beta = 1$ 时 $z$ 就是一坨噪声,$\text{std}_x(\mu)$ 只剩 0.003。

那么 scaling_factor 是什么?在这里它是原始 latent 标准差的倒数。LDM 的 原始标定代码 使用 1. / z.flatten().std();diffusers 文档中的标准差描述要结合实际乘法方向理解。本文 toy 的跨维 RMS 也不是一般情况下与展平标准差完全相等:后者还包含不同通道均值的差异。

$$z_{\text{scaled}} = s \cdot z, \qquad s = \frac{1}{\text{std}(z)}$$

Stable Diffusion 取 $s = 0.18215$,倒过来就是 $\text{std}(z) = 5.4900$、$\text{Var}(z) = 30.1399$。为什么需要这一步,LDM D.1 用信噪比的语言回答了:扩散模型的噪声表 $\sigma_t$ 是相对数据尺度标定的,把这个比值写出来

$$\text{SNR} = \frac{\text{Var}(z)\,\bar{\alpha}}{1 - \bar{\alpha}}$$

模型以为的是 $\bar{\alpha}/(1-\bar{\alpha})$,实际却差 $\text{Var}(z) = 30.14$ 倍,也就是 14.79 dB。每一步模型都以为「噪声占了这么多」,实际噪声只有它以为的 1/30。

最后一个式子把「差 30 倍」翻译成「图发灰」。假设扩散模型是在单位方差潜空间上训好的,那它的去噪器就是一个高斯 MMSE 估计:先验 $z_{0} \sim \mathcal{N}(0, \sigma_{p}^{2} I)$,观测 $z_{t} = \sqrt{\bar{\alpha}}\,z_{0} + \sqrt{1-\bar{\alpha}}\,\epsilon$,后验均值是

$$\hat{z}_{0} = \frac{\sigma_{p}^{2}\sqrt{\bar{\alpha}}}{\sigma_{p}^{2}\bar{\alpha} + 1 - \bar{\alpha}}\, z_{t}$$

代入 $\sigma_{p}^{2} = 1$ 得 $\hat{z}_{0} = \sqrt{\bar{\alpha}}\, z_{t}$;而真实潜空间的标准差是 $\sigma$,正确的 MMSE 估计要用 $\sigma_{p}^{2} = \sigma^{2}$。两者相除就是第 01 节那个 $r(\bar{\alpha}) = \bar{\alpha} + (1-\bar{\alpha})/\sigma^{2}$。$\sigma = 5.49$ 时它在 $\bar{\alpha} \to 0$ 处趋于 $1/\sigma^{2} = 0.033$——重建幅度只剩 3.3%,图当然是灰的。

04. 代码实现

完整脚本在文末附录(vae_minimal.py、collapse_closed_form.py、latent_scaling.py、make_figures.py),只依赖 numpy 与 matplotlib,全部用 /usr/local/bin/python3 实跑过,下面每个数字都是真实输出。环境里没有 torch,所以反向传播是手推的——这反而更好,公式和代码能一行行对上。

4.1 最小 VAE 与两轮对照实验

模型就是 3.1 节那套,符号一一对应:

def forward(self, x, eps):
    p = self.p
    h1 = np.maximum(x @ p["W1"] + p["b1"], 0.0)
    mu = h1 @ p["Wmu"] + p["bmu"]                     # 后验均值
    lv_raw = h1 @ p["Wlv"] + p["blv"]
    lv = np.clip(lv_raw, LOGVAR_MIN, LOGVAR_MAX)      # logvar 截断,防 exp 溢出
    sig = np.exp(0.5 * lv)                            # 后验标准差 sigma
    z = mu + sig * eps                                # 重参数化
    g1 = np.maximum(z @ p["V1"] + p["c1"], 0.0)
    xhat = g1 @ p["V2"] + p["c2"]
    return xhat, dict(x=x, h1=h1, mu=mu, lv=lv, lv_raw=lv_raw,
                      sig=sig, z=z, g1=g1, xhat=xhat, eps=eps)

@staticmethod
def losses(x, xhat, mu, lv, sig):
    recon = float(np.mean((xhat - x) ** 2))                        # 重建:MSE
    kl_dim = 0.5 * np.mean(mu ** 2 + sig ** 2 - lv - 1.0, axis=0)  # 逐维 KL(3.2 节那个式子)
    return recon, kl_dim, float(np.sum(kl_dim))                    # 对维度求和

数据是合成的:64 个观测维度由 4 个高斯因子线性混合而成,再逐维归一化到方差 1(所以「重建 MSE = 1」= 什么都没学到)。潜变量给了 8 维,故意比真实因子数多一倍,看模型怎么选。编码器/解码器各一层 128 宽的隐藏层,Adam,250 epoch,逐维 KL 大于 0.05 nat 才算「存活」。

跑两轮:一组不加重衰减,一组给权重加上 $10^{-3}$ 的衰减。结果:

    weight_decay      beta    重建MSE        总KL    std(z)    1/std   存活维度
--------------------------------------------------------------------------------
               0     1e-06     0.0030     76.958    2.0596   0.4855      8/8
               0     1e-05     0.0031     39.861    1.4107   0.7088      8/8
               0     0.0001     0.0038     23.512    1.0651   0.9389      7/8
               0      0.001     0.0062     12.454    1.0067   0.9934      4/8
               0       0.01     0.0247      7.843    1.0055   0.9946      4/8
               0        0.1     0.2119      3.092    0.9952   1.0048      4/8
               0        0.3     0.6197      0.887    1.0000   1.0000      3/8
               0          1     1.0000      0.000    1.0001   0.9999      0/8
               0          3     1.0001      0.000    1.0000   1.0000      0/8
           0.001     1e-06     0.0046     55.983    2.7174   0.3680      8/8
           0.001     1e-05     0.0046     49.327    2.5978   0.3849      6/8
           0.001     0.0001     0.0049     28.672    2.0580   0.4859      5/8
           0.001      0.001     0.0075     13.409    1.3082   0.7644      4/8
           0.001       0.01     0.0259      7.917    1.0689   0.9355      4/8
           0.001        0.1     0.2148      3.107    1.0051   0.9949      4/8
           0.001        0.3     0.6301      0.858    1.0018   0.9982      3/8
           0.001          1     1.0000      0.000    1.0000   1.0000      0/8
           0.001          3     1.0000      0.000    1.0000   1.0000      0/8

三件事一眼可见:

第一,坍缩是一条断崖,不是一个缓坡。 $\beta$ 从 $10^{-2}$ 到 $0.1$ 到 $0.3$ 到 $1$,重建 MSE 走 $0.0247 \to 0.2119 \to 0.6197 \to 1.0000$。活着的维度 $4 \to 4 \to 3 \to 0$。$\beta = 1$ 时总 KL 精确变成 0.000,说明近似后验和先验完全重合——编码器变成了一个只会输出 $\mathcal{N}(0,I)$ 的函数。

第二,维度是一个个死的。 $\beta = 10^{-3}$ 时逐维 KL 是 [2.96, 0.002, 3.19, 0.004, 3.12, 3.18, 0.001, 0.001]——恰好 4 个在 3 nats 附近,另外 4 个趴在 0.002 上。活下来的正好是 4 个,和合成数据的真实因子数相等。 模型自己算出「值得买 4 个维度」,这不是我告诉它的。

第三,潜空间尺度确实跟着 $\beta$ 漂。 无权重衰减时从 1.0067($\beta=10^{-3}$)漂到 2.0596($\beta = 10^{-6}$);加上 $10^{-3}$ 的权重衰减后在同样区间漂到 2.7174。对应的 scaling_factor 从 0.9934 掉到 0.4855 / 0.3680。这就是「0.18215 从哪来」的机制:$\beta$ 小到 KL 项失去约束力时,潜空间的尺度不再由任何东西钉住。

需要说清楚的是,让尺度长大的那股力在我的实验里是权重衰减(潜尺度越大,解码器权重就可以越小),而 Adam 自身对参数尺度不敏感/敏感的部分也在推它(不加重衰减时也会从 1.0067 漂到 2.0596)。真实 VAE 里起同样作用的还有编码器末端的归一化层、初始化尺度、以及训练超参。结论只需要一条:$\beta$ 决定的是「KL 有没有能力把尺度钉在 1」,钉不住之后具体漂到几,是别的因素决定的。 Stable Diffusion 漂到了 5.49,我的玩具漂到了 2.7,同一个机制。

4.2 换条路验一遍:一维闭式解

3.3 节那组闭式解值得单独验,因为「后验坍缩阈值 $\lambda = 1$」这个结论如果错了,整篇文章的框架就错了。做法是直接对 $(w, m, \log s)$ 做梯度下降。

def closed_form(beta, v=V, sigma_x=SIGMA_X):
    lam = beta * sigma_x ** 2 / v
    if lam >= 1.0:                  # lambda >= 1:坍缩
        return dict(lam=lam, u=0.0, s2=1.0, w2=0.0, collapsed=True)
    u = 1.0 - lam
    return dict(lam=lam, u=u, s2=lam, w2=v * u, collapsed=False)

($v = 1$、$\sigma_{\text{dec}} = 1$,所以 $\lambda = \beta$、坍缩阈值 $\beta^{*} = 1$。)对照结果:

  beta   lambda |    u=wm 预测        拟合 |    s^2 预测        拟合 |    w^2 预测        拟合 | 坍缩
  0.05    0.050 |     0.9500    0.9500 |    0.0500    0.0500 |    0.9500    0.9500 | 否
   0.1    0.100 |     0.9000    0.9000 |    0.1000    0.1000 |    0.9000    0.9000 | 否
   0.2    0.200 |     0.8000    0.8000 |    0.2000    0.2000 |    0.8000    0.8000 | 否
   0.4    0.400 |     0.6000    0.6000 |    0.4000    0.4000 |    0.6000    0.6000 | 否
   0.6    0.600 |     0.4000    0.4000 |    0.6000    0.6000 |    0.4000    0.4000 | 否
   0.8    0.800 |     0.2000    0.2000 |    0.8000    0.8000 |    0.2000    0.2000 | 否
  0.95    0.950 |     0.0500    0.0500 |    0.9500    0.9500 |    0.0500    0.0500 | 否
     1    1.000 |     0.0000    0.0000 |    1.0000    1.0000 |    0.0000    0.0000 | 是
   1.2    1.200 |     0.0000    0.0000 |    1.0000    1.0000 |    0.0000    0.0000 | 是
     2    2.000 |     0.0000    0.0000 |    1.0000    1.0000 |    0.0000    0.0000 | 是
     5    5.000 |     0.0000    0.0000 |    1.0000    1.0000 |    0.0000    0.0000 | 是

未坍缩区间内,闭式解与数值拟合的最大偏差:5.91e-06

坍缩点 (w, m, s) = (0, 0, 1) 处的梯度(应当恒为 0):
  beta=0.05   dJ/dw=+0.00e+00  dJ/dm=+0.00e+00  dJ/ds=+0.00e+00
  beta=0.5    dJ/dw=+0.00e+00  dJ/dm=+0.00e+00  dJ/ds=+0.00e+00
  beta=1      dJ/dw=+0.00e+00  dJ/dm=+0.00e+00  dJ/ds=+0.00e+00
  beta=5      dJ/dw=+0.00e+00  dJ/dm=+0.00e+00  dJ/ds=+0.00e+00

坍缩点 vs 解析解的目标函数值(谁小谁是最优):
  beta=0.05   J(解析解)=  0.09989  J(坍缩点)=  0.50000
  beta=0.5    J(解析解)=  0.42329  J(坍缩点)=  0.50000
  beta=0.9    J(解析解)=  0.49741  J(坍缩点)=  0.50000
  beta=1      J(解析解)=  0.50000  J(坍缩点)=        —
  beta=2      J(解析解)=  0.50000  J(坍缩点)=        —

三处细节值得留意。偏差 5.91e-06 说明闭式解是对的。梯度恒等于 0 说明坍缩点在任意 $\beta$ 下都是驻点——它一直「在那儿」,$\lambda \ge 1$ 时它只是终于变成了最小值。$\beta = 0.9$ 时两者的目标值只差 0.0026,说明接近阈值时塌向坍缩点的阻力非常小,这解释了为什么真实训练里坍缩一旦开始就很快。

4.3 后验坍缩长什么样

这张图要看什么:三张子图是三档 $\beta$ 下逐维 KL 的柱状图,蓝色是存活维度(KL > 0.05 nat),灰色是死掉的。注意三张图的纵轴量级完全不同(3.19 / 0.39 / 0.00001 nats,差五个数量级),所以我把每张图各自缩放并标了纵轴最大值——如果共享纵轴,后两张会被压成一条线,看不出结构。$\beta = 10^{-3}$ 时是 4 根高柱加 4 根贴地;$\beta = 0.3$ 时只剩 3 根矮柱,重建 MSE 已经涨到 0.6197;$\beta = 1$ 时一根都没有。

4.4 潜尺度漂移:从 1.0067 到 2.7174

这张图左侧按全方差公式堆叠两项方差:跨样本后验均值方差、平均后验方差;右侧显示第一项占总方差的比例。两项方差相加后开方才得到边际标准差,标准差本身不能直接堆叠。这个分解描述二阶统计,不能等同于信息与噪声的互信息分解。

这张图帮助理解尺度与信息的区别:相同的边际方差可以来自不同的均值/条件方差组合。缩放只调整统计尺度,不保证 latent 的语义分布匹配,也不能单凭 std 判断是否坍缩。

4.5 忘掉 scaling_factor 的代价,量化

这张图要看什么:左图的纵轴是对数的,三条线分别是三种潜空间尺度下的去噪幅度比 $r$。绿线 $\sigma = 1.0$ 平在 $r = 1$ 上(正确);蓝线是真实情形 $\sigma = 5.49$,$\bar{\alpha}$ 越小掉得越狠,最左端贴在 $1/\sigma^{2} = 0.033$,也就是幅度只剩 3.3%(后验均值幅度偏小);红线是反向错误——潜空间其实是单位方差却照抄了 0.18215,$\sigma$ 变成 0.182,$r$ 冲到 30 倍(后验均值幅度偏大)。右图是同一个前向过程下两种去噪器的均方误差,中段拉开 8 倍,两端收敛($\bar{\alpha} \to 1$ 时都没噪声要除,$\bar{\alpha} \to 0$ 时都没信息可用——误差差距最大的地方在中间,这跟直觉不太一样)。

   alpha_bar |      错假设 MSE       正确 MSE         理论后验方差       恶化倍数
       0.999 |       0.0010       0.0010         0.0010       1.03x
        0.99 |       0.0130       0.0101         0.0101       1.28x
         0.9 |       0.3933       0.1106         0.1107       3.55x
         0.5 |       7.8154       0.9682         0.9679       8.07x
         0.1 |      24.5816       6.9502         6.9305       3.54x
        0.01 |      29.6359      23.2187        23.1056       1.28x

「理论后验方差」一列是和蒙特卡洛结果并排校核的:$\bar{\alpha}=0.5$ 处 0.9679 对 0.9682,这一格的差值约 $3\times10^{-4}$;其他格的采样波动会更大,不能把单个差值当作整张表的误差保证。

05. 工业级实现对照

真实框架长什么样,看 diffusers 的 AutoencoderKL 和 DiagonalGaussianDistribution。以下以 2026-09 时的实现为准,上游会重构。

编码链路(AutoencoderKL._encode / encode):encoder(x) → quant_conv → 把结果塞进 DiagonalGaussianDistribution。和最小实现的差异有四处:

  • 用卷积而不是全连接。最小实现里 x 是一个 64 维向量,一层矩阵乘就够;图像要保留空间结构,所以编码器是卷积堆栈,输出 [B, 2*z_channels, H/8, W/8],逐像素各出一组 $(\mu, \text{logvar})$。KL 也因此对每个元素都算一次。
  • *quant_conv 是 1×1 卷积,默认从 `2latent_channels映到同样的通道数**。编码器已经输出均值与 logvar 所需的双倍通道;这里做通道混合,随后torch.chunk(parameters, 2, dim=1)` 分成两组,不是在 quant_conv 这一步翻倍。
  • logvar 被硬截断到 $[-30, 20]$,这一行是 self.logvar = torch.clamp(self.logvar, -30.0, 20.0)。别小看它:$\exp(20) = 4.85\times10^{8}$,$\exp(-30) = 9.36\times10^{-14}$。我的最小实现里也照抄了这个阈值。不加截断,一个异常值就能让 KL 炸到 1e8 或者把梯度打进下溢区。
  • encode 返回的是后验分布,不乘 scaling factor。这一步由调用方负责——pipeline 里显式写 latents = vae.encode(image).latent_dist.sample() * vae.config.scaling_factor。这个设计是个典型的「容易漏」:它不在 encode 里面,所以复制粘贴半段代码就会丢掉。

推理时用 sample() 还是 mode():两者都合法,必须与具体管线的约定一致。mode() 返回均值,适合需要确定性编码的实验;sample(generator=...) 返回后验样本,固定随机生成器也能复现。diffusers 的 Stable Diffusion img2img 中 retrieve_latents 默认走 sample,所以不能说图生图、修补必须用 mode()。deterministic=True 是把分布对象的方差与标准差置零的特殊模式,不等于重新训练过的普通自编码器。

KL 的实现细节:kl() 里写的是 0.5 * torch.sum(torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar, dim=[1, 2, 3])。跟 3.2 节那个式子逐字对上(var 就是 $\sigma^{2}$,logvar 就是 $\log\sigma^{2}$),只是求和范围从「8 个潜变量维度」变成「$4 \times 64 \times 64 = 16384$ 个元素」,返回的是每个样本一个标量。这一点是全文最容易被忽略的工程事实:KL 是逐元素求和的,所以潜变量个数一变,同一个 $\beta$ 数值的含义完全不同。LDM 用 $10^{-6}$、我的实验用 $10^{-3}\sim1$,这两个数根本不在同一个坐标系里——看别人的 $\beta$ 必须连着看它的潜变量张量形状。

视频 VAE 上这条线怎么延伸:时间维一起下采样,潜变量变成 [B, C, T/4, H/8, W/8],KL 求和范围又大了一个量级。KL 虽逐元素求和,但编码器、解码器会耦合各维,不能保证各维独立开关。元素数增大时还要同时看重建项如何归约,不能仅凭维度数断言同一 β 必然关闭更多维度——这也是为什么视频 VAE 的 loss 配比需要单开一篇(见 视频 VAE 的常见 loss 组合)。

带 shift 的 VAE 要区分编码和解码方向:例如 diffusers 的 SD3 管线 在解码前使用 $z_{\text{raw}}=z_{\text{diffusion}}/\text{scale}+\text{shift}$;对应的正向变换才是 $z_{\text{diffusion}}=(z_{\text{raw}}-\text{shift})\cdot\text{scale}$。读取具体 checkpoint 的配置和管线,不把 SD1.x 的常数套到 SD3 / FLUX。

06. 代价与边界

VAE 是有损压缩,这是第一位的代价。 以 $f=8$、4 通道为例,一张 512×512 的图进来,出去的是 $4 \times 64 \times 64 = 16384$ 个数,压缩比 $3\times512\times512 / 16384 \approx 48$。压缩本身就是有损的,而且丢的是高频——文字、小脸、细纹理这些恰恰是人类最敏感的东西。扩散模型再强也补不回来,因为信息在进入扩散过程之前就已经没了(这也是 SD 生态里独立高分辨率精修模型存在的理由)。

把 $\beta$ 调小,赔的是潜空间的可预测性。 潜空间尺度失去约束之后:换 VAE 必须核对 latent 语义、尺度与扩散模型训练约定(不能只重算一个系数),而且——按 LDM D.1 的观察——即使一致地训练,信噪比全程偏高也会让高分辨率卷积采样出问题。$\beta$ 越小,重建越好,但潜空间越像一个「定制格式」,越难被别的东西复用。

把 $\beta$ 调大,赔的是潜变量的信息容量。 实测 $\beta=1$ 时 8 个维度接近先验、重建 MSE 约为 1。在线性闭式模型里,存活区的 MSE 为 $\beta\sigma_{\text{dec}}^2$,到阈值后连续接到 $v$;逐维 KL 也连续降到 0。采用阈值统计的“存活维度数”会出现台阶,但这不等于重建误差存在不连续跳变。

什么时候不该用潜空间:需要像素级保真的任务(超分、医学影像、文字密集的文档生成)直接上像素空间或多尺度方案,别压 $f=8$;另外如果任务的训练数据量和算力都充足、又不要求高分辨率,2.1 节那套「潜空间省算力」的收益应与压缩误差一起实测。

另一条路线:VQ。 用有限码本替换连续高斯,带来量化误差、码本容量与使用率等另一组权衡。LDM 论文观察到它所用的 VQ latent 方差接近 1;这不是离散化的数学保证,码本仍然可以整体缩放,是否需要归一化要看实际训练分布。

07. 经典论文脉络

  1. Kingma & Welling, 2013, Auto-Encoding Variational Bayes(arXiv:1312.6114) — VAE 本体。两个贡献撑起了后面所有工作:重参数化技巧(把采样挪出计算图)和对角高斯先验下 KL 的闭式解(3.2 节那个式子)。
  2. Higgins 等, ICLR 2017, β-VAE: Learning Basic Visual Concepts with a Constrained Variational Framework — 把 KL 权重从固定的 1 变成可调旋钮,$\beta>1$ 换解耦表征。副产品是把「后验坍缩」推到了台前:$\beta$ 一大,重建就塌。
  3. van den Oord 等, 2017, Neural Discrete Representation Learning(arXiv:1711.00937) — VQ-VAE。用码本查表替换连续高斯采样,潜空间变成离散索引,3.2 节那个 KL 项整个消失了。
  4. Esser 等, CVPR 2021, Taming Transformers for High-Resolution Image Synthesis(arXiv:2012.09841) — VQGAN。给 VQ 自编码器加上感知损失和对抗损失,把重建质量推到可商用,第一次让「学到的潜空间 + 自回归」在 1024² 上真正可用。
  5. Rombach 等, CVPR 2022, High-Resolution Image Synthesis with Latent Diffusion Models(arXiv:2112.10752) — LDM / Stable Diffusion。它把「KL 权重、潜空间方差、信噪比、scaling factor」这四件事的关系写进了 4.3.2 与 D.1 两节,$0.18215$ 从此钉在了所有下游代码里。

顺着这条线看,故事其实是「潜空间的统计性质从哪儿来」在被一步步讲清楚:2013 年给了目标函数,2017 年发现权重会毁掉它,2017—2021 年绕开它(离散化 + 对抗训练),2022 年终于正面处理它的尺度问题。

08. 常见误解

误解一:KL 既然是正则项,越大越好。 KL 确实可以看作正则,也可用信息预算解释,但增强它会牺牲重建。第 09 节的密集扫描显示 MSE 从 $0.0650$ 到 $0.8874$ 逐步上升;台阶出现在人为设阈值的存活维度计数,不能把稀疏扫描误读为“中间没有过渡”。

误解二:任意 KL 权重下都在拟合原模型的真实后验。 当 $\beta=1$ 且变分族足够时,最优 $q$ 可以等于该生成模型的真实后验;本文线性高斯模型就是例子。$\beta\ne1$ 则改变权衡,存活区 $s^2=\beta\sigma_{\text{dec}}^2/v$ 同时依赖权重、观测噪声和数据方差 $v$,不能说与数据无关。

误解三:聚合后验等于先验就意味着坍缩。 逐样本 $q(z\mid x)$ 和聚合后验 $q(z)=\int q(z\mid x)p_{\text{data}}(x)\,dx$ 是两个对象。本文未坍缩线性解满足 $m^2v+s^2=(1-\lambda)+\lambda=1$,因而 $q(z)=\mathcal N(0,1)$,但 $m\ne0$,仍然传递信息。真正的完全坍缩要求几乎所有输入的整个 $q(z\mid x)$ 都等于同一个先验。

误解四:scaling factor 是固定的魔法常数。 LDM 的标定实现使用首个训练 batch 的 latent 展平后的标准差,令缩放因子为其倒数;这不是逐通道独立归一化。使用预训练模型时应遵守 checkpoint 保存的系数,不能随手在一张新图上重估并替换。换 VAE 还可能改变 latent 的语义和通道分布,重算一个标量不保证与原扩散模型兼容。

误解五:后验采样只在训练时用。 官方 img2img 管线默认也会采样;固定 generator 可以控制随机性。mode() 是去掉编码采样噪声的一种选择,必须核对训练与推理的分布约定,不能统一替换所有管线。

误解六:截断范围等于所有精度下的安全范围。 [-30,20] 是实现中的防护范围,但 exp(20) 超过 fp16 最大有限值,实际还要看计算 dtype 和 VAE 是否上转 fp32。实验中的 logvar 约 −9.4 离 −30 很远,不能仅凭这个数断言已接近数值崩溃。

09. 动手验证

三个小实验,都能在两分钟内跑完。

实验一:确认坍缩是「跳」还是「滑」。 一行命令:

python vae_minimal.py --betas=0.03,0.05,0.1,0.2,0.5 --wd=0

实测($\text{weight\_decay} = 0$,$250$ epoch,随机种子固定,可复现):

$\beta$ $0.03$ $0.05$ $0.1$ $0.2$ $0.5$
重建 MSE $0.0650$ $0.1057$ $0.2119$ $0.4203$ $0.8874$
总 KL $5.580$ $4.554$ $3.092$ $1.676$ $0.198$
存活维度 $4/8$ $4/8$ $4/8$ $4/8$ $1/8$

结论:重建误差逐渐增加,存活维度的计数呈台阶。 这里把每维 KL 超过 0.05 nat 计为存活,所以计数天然离散。在线性模型中,每维 KL 为 $-\tfrac12\log\lambda$(未坍缩时),连续趋向 0,并不会从 3 直接跳到 0;应同时观察逐维 KL、重建与均值/方差统计。

实验二:确认尺度漂移是权重衰减带来的。 把权重衰减单独开到 $10^{-2}$:

python vae_minimal.py --wd=0.01

三档权重衰减下 $\text{std}(z)$ 的对照(同一批 $\beta$,越小越说明潜尺度已经失控):

$\beta$ $\text{wd} = 0$ $\text{wd} = 10^{-3}$ $\text{wd} = 10^{-2}$
$10^{-6}$ $2.0596$ $2.7174$ $2.6749$
$10^{-5}$ $1.4107$ $2.5978$ $2.6383$
$10^{-4}$ $1.0651$ $2.0580$ $2.5399$
$10^{-3}$ $1.0067$ $1.3082$ $2.0463$
$10^{-2}$ $1.0055$ $1.0689$ $1.3657$
$10^{-1}$ $0.9952$ $1.0051$ $1.0523$
$3\times10^{-1}$ $1.0000$ $1.0018$ $1.0000$

看这张表的方式是看「回到 1.00 的那个拐点」在往右挪:$\text{wd} = 0$ 时 $\beta \ge 10^{-3}$ 就已经归位;$\text{wd} = 10^{-3}$ 时要到 $\beta \ge 10^{-2}$;$\text{wd} = 10^{-2}$ 时要一路推到 $\beta \ge 0.3$。权重衰减把「潜尺度失控」的区间往大 $\beta$ 方向整整推了两个数量级。一个诚实的补充:$10^{-6}$ 那一档 $10^{-2}$ 的 $2.6749$ 反而略低于 $10^{-3}$ 的 $2.7174$,说明这个偏移会饱和,不是权重衰减越大越离谱。

实验三:反向验证 scaling factor。 换两个反事实的缩放系数各跑一次(--sf 会把整张表按新系数重算):

python latent_scaling.py --sf=0.5
python latent_scaling.py --sf=2.0

实际结果:

$\text{scaling factor}$ $\text{std}(z)$ $r$ 在 $\bar{\alpha}=0.999$ $r$ 在 $\bar{\alpha}=0.01$ $\bar{\alpha}=0.5$ 处 MSE 比值
$0.18215$(真实值) $5.4900$ $0.9990$ $0.0428$ $8.07\times$
$0.5$(缩放不足) $2.0000$ $0.9992$ $0.2575$ $1.57\times$
$2.0$(缩放过头) $0.5000$ $1.0030$ $3.9700$ $1.56\times$

$\sigma = 2$ 时 $r$ 全程在 1 以下(最低 $0.2575$,正是 $1/\sigma^{2} = 0.25$ 加上 $\bar{\alpha}$ 那一项),该高斯估计器的幅度偏小;$\sigma = 0.5$ 时 $r$ 全程在 1 以上(最高 $3.97$,趋近 $1/\sigma^{2} = 4$),该高斯估计器的幅度偏大。在这个高斯模型中:$r$ 是大于还是小于 1,只取决于 $\sigma$ 是大于还是小于 1;而且 $\bar{\alpha}\to0$ 的末端,$r$ 就直接收敛到 $1/\sigma^{2}$ —— 噪声越大,缩放错误的代价越彻底。反过来看真实值那一档:$r$ 掉到 $0.0428$,MSE 差 $8.07$ 倍,比两个假想档位严重一个量级,这就是 0.18215 这个数不能省的原因。

10. 延伸阅读

附录:完整代码

09 节用到的脚本全文如下(vae_minimal.py、collapse_closed_form.py、latent_scaling.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。

vae_minimal.py

"""最小 VAE:numpy 手写反向传播,扫 KL 权重 beta,观察潜空间统计量怎么变。

这篇文章要回答的问题是:KL 权重 beta 调大调小,到底改变了什么?
本脚本做两轮对照实验:

  条件 A(无权重衰减):潜尺度被 KL 项钉住,beta 从小到大,std(z) 稳在 1 附近。
  条件 B(权重衰减 1e-3):解码器偏好「潜尺度大、权重小」的解,
                          只有 KL 项拦得住它;beta 一小,潜尺度就漂走。

两轮对照说明同一件事:**beta 并不直接决定潜尺度,它决定的是
「KL 项有没有能力把潜尺度钉在 1」**。钉不住时,潜尺度由训练里其他所有力
(权重衰减、归一化、初始化)共同决定,可以漂到 5 倍开外——
此实验解释弱正则下的一种漂移机制,但不能据此反推 Stable Diffusion 的具体训练轨迹。

运行:  /usr/local/bin/python3 vae_minimal.py
依赖:  numpy
"""

import os
import json
import numpy as np

# ── 固定随机性,保证正文里贴的每个数字都能复现 ─────────────────────────
DATA_SEED = 0
INIT_SEED = 1
BATCH_SEED = 2
EVAL_SEED = 3

N_TRAIN = 4096      # 训练样本数
D_OBS = 64          # 观测维度(类比一张图的像素数)
H_HID = 128         # 编码器/解码器隐藏层
D_LAT = 8           # 潜变量维度
K_FAC = 4           # 合成数据的真实因子数(< D_LAT,故意留出冗余维度)
EPOCHS = 250
BATCH = 256
LR = 3e-3
LOGVAR_MIN, LOGVAR_MAX = -30.0, 20.0   # 与 diffusers DiagonalGaussianDistribution 一致

BETAS = [1e-6, 1e-5, 1e-4, 1e-3, 1e-2, 1e-1, 3e-1, 1.0, 3.0]
WD_LIST = [0.0, 1e-3]                  # 两轮对照:无权重衰减 / 有权重衰减


# ─────────────────────────── 合成数据 ───────────────────────────
def make_data(n=N_TRAIN, d=D_OBS, k=K_FAC, seed=DATA_SEED):
    """x = A u + 观测噪声,u 是 k 维高斯因子。

    做法和真实图像一样:观测维度很多(d=64),但真正驱动它的因子只有 k=4 个。
    潜变量维度 D_LAT=8 > 4,所以模型必须自己决定「用几个维度」,
    这正是后验坍缩能被观察到的前提。
    最后把每个观测维度归一化到方差 1,这样「重建 MSE = 1」就等于「什么都没学到」。
    """
    rng = np.random.default_rng(seed)
    A = rng.normal(size=(d, k)) / np.sqrt(k)
    U = rng.normal(size=(n, k))
    X = U @ A.T + 0.05 * rng.normal(size=(n, d))
    X = X / X.std(axis=0, keepdims=True)      # 逐维方差 = 1
    return X


# ─────────────────────────── 模型 ───────────────────────────
class VAE:
    """编码器 x -> h -> (mu, logvar),解码器 z -> g -> x_hat。

    符号与正文第 03 节一致:
        h      = relu(x W1 + b1)              编码器隐藏层
        mu     = h Wmu + bmu                  后验均值
        logvar = h Wlv + blv                  后验对数方差(网络出的是 log sigma^2)
        z      = mu + exp(0.5 logvar) * eps   重参数化
        x_hat  = relu(z V1 + c1) V2 + c2      解码器
    """

    def __init__(self, d_obs=D_OBS, h=H_HID, d_lat=D_LAT, seed=INIT_SEED):
        rng = np.random.default_rng(seed)
        sc = lambda fan_in, fan_out: rng.normal(
            scale=np.sqrt(2.0 / (fan_in + fan_out)), size=(fan_in, fan_out)
        )
        self.p = {}
        self.p["W1"] = sc(d_obs, h)
        self.p["b1"] = np.zeros(h)
        self.p["Wmu"] = sc(h, d_lat)
        self.p["bmu"] = np.zeros(d_lat)
        self.p["Wlv"] = sc(h, d_lat)
        self.p["blv"] = np.zeros(d_lat)
        self.p["V1"] = sc(d_lat, h)
        self.p["c1"] = np.zeros(h)
        self.p["V2"] = sc(h, d_obs)
        self.p["c2"] = np.zeros(d_obs)
        self.d_obs, self.h, self.d_lat = d_obs, h, d_lat

    def forward(self, x, eps):
        p = self.p
        h1 = np.maximum(x @ p["W1"] + p["b1"], 0.0)
        mu = h1 @ p["Wmu"] + p["bmu"]
        lv_raw = h1 @ p["Wlv"] + p["blv"]
        lv = np.clip(lv_raw, LOGVAR_MIN, LOGVAR_MAX)   # 防 exp 溢出,见第 05 节
        sig = np.exp(0.5 * lv)
        z = mu + sig * eps
        g1 = np.maximum(z @ p["V1"] + p["c1"], 0.0)
        xhat = g1 @ p["V2"] + p["c2"]
        cache = dict(x=x, h1=h1, mu=mu, lv=lv, lv_raw=lv_raw, sig=sig,
                     z=z, g1=g1, xhat=xhat, eps=eps)
        return xhat, cache

    @staticmethod
    def losses(x, xhat, mu, lv, sig):
        """recon = 逐元素均方误差;kl = 0.5 * sum_d(mu^2 + var - 1 - logvar)。

        注意 kl 是对潜变量维度「求和」而不是求平均——这是标准写法,
        也意味着 beta 的等效大小会随潜变量个数一起变。
        """
        recon = float(np.mean((xhat - x) ** 2))
        kl_dim = 0.5 * np.mean(mu ** 2 + sig ** 2 - lv - 1.0, axis=0)  # [d]
        return recon, kl_dim, float(np.sum(kl_dim))

    def backward(self, cache, beta):
        """全部手推。grad 的形状与 self.p 一一对应。"""
        p = self.p
        x, h1, mu, lv, lv_raw, sig, z, g1, xhat, eps = (
            cache[k] for k in ["x", "h1", "mu", "lv", "lv_raw", "sig", "z", "g1", "xhat", "eps"]
        )
        B, D = x.shape[0], self.d_obs
        g = {k: np.zeros_like(v) for k, v in p.items()}

        # 重建项:recon = mean((xhat - x)^2),d/d(xhat) = 2 (xhat - x) / (B*D)
        dxhat = 2.0 * (xhat - x) / (B * D)
        g["V2"] += g1.T @ dxhat
        g["c2"] += dxhat.sum(axis=0)
        dg1 = (dxhat @ p["V2"].T) * (g1 > 0)          # relu
        g["V1"] += z.T @ dg1
        g["c1"] += dg1.sum(axis=0)
        dz = dg1 @ p["V1"].T                          # [B, d]

        # 重参数化:dz 同时流回 mu 与 logvar 两条支路
        dmu = dz + beta * mu / B                      # KL 对 mu 的导数 = beta * mu
        dlv = dz * (0.5 * sig * eps) + beta * 0.5 * (sig ** 2 - 1.0) / B
        # 被 clamp 的位置梯度必须截断,否则会用未截断的 logvar 继续更新
        dlv = dlv * ((lv_raw > LOGVAR_MIN) & (lv_raw < LOGVAR_MAX))

        g["Wmu"] += h1.T @ dmu
        g["bmu"] += dmu.sum(axis=0)
        g["Wlv"] += h1.T @ dlv
        g["blv"] += dlv.sum(axis=0)
        dh1 = (dmu @ p["Wmu"].T + dlv @ p["Wlv"].T) * (h1 > 0)
        g["W1"] += x.T @ dh1
        g["b1"] += dh1.sum(axis=0)
        return g


# ─────────────────────────── Adam ───────────────────────────
class Adam:
    def __init__(self, params, lr=LR, b1=0.9, b2=0.999, eps=1e-8):
        self.m = {k: np.zeros_like(v) for k, v in params.items()}
        self.v = {k: np.zeros_like(v) for k, v in params.items()}
        self.lr, self.b1, self.b2, self.eps = lr, b1, b2, eps
        self.t = 0

    def step(self, params, grads):
        self.t += 1
        for k in params:
            self.m[k] = self.b1 * self.m[k] + (1 - self.b1) * grads[k]
            self.v[k] = self.b2 * self.v[k] + (1 - self.b2) * grads[k] ** 2
            mhat = self.m[k] / (1 - self.b1 ** self.t)
            vhat = self.v[k] / (1 - self.b2 ** self.t)
            params[k] -= self.lr * mhat / (np.sqrt(vhat) + self.eps)


def train(beta, X, weight_decay=0.0, epochs=EPOCHS, batch=BATCH, seed=BATCH_SEED):
    """训练一个 VAE。weight_decay 是本次实验的关键对照变量。"""
    rng = np.random.default_rng(seed)
    model = VAE()
    opt = Adam(model.p)
    n = X.shape[0]
    for _ in range(epochs):
        perm = rng.permutation(n)
        for s in range(0, n, batch):
            xb = X[perm[s:s + batch]]
            eps = rng.normal(size=(xb.shape[0], D_LAT))
            xhat, cache = model.forward(xb, eps)
            model.losses(xb, xhat, cache["mu"], cache["lv"], cache["sig"])
            grads = model.backward(cache, beta)
            if weight_decay > 0:                      # 只对权重做衰减,偏置不管
                for k in ["W1", "Wmu", "Wlv", "V1", "V2"]:
                    grads[k] += weight_decay * model.p[k]
            opt.step(model.p, grads)
    return model


def evaluate(model, X, beta, weight_decay=0.0, n_repeat=8):
    """统计潜空间性质,并把边际标准差拆成两半:

        Var(z_j) = Var_x(mu_j(x)) + E_x[sigma_j(x)^2]

    前一半是「不同样本被编码到不同位置」,后一半是「每个样本自己抖多少」。
    scaling_factor 要归一化的是两者之和的开方。
    """
    rng = np.random.default_rng(EVAL_SEED)
    n, d = X.shape[0], D_LAT
    mu_all = np.empty((n, d))
    var_all = np.empty((n, d))
    recon_all = 0.0
    kl_dim_all = np.zeros(d)

    for s in range(0, n, 512):
        xb = X[s:s + 512]
        B = xb.shape[0]
        h1 = np.maximum(xb @ model.p["W1"] + model.p["b1"], 0.0)
        mu = h1 @ model.p["Wmu"] + model.p["bmu"]
        lv = np.clip(h1 @ model.p["Wlv"] + model.p["blv"], LOGVAR_MIN, LOGVAR_MAX)
        var = np.exp(lv)
        sig = np.sqrt(var)
        mu_all[s:s + B] = mu
        var_all[s:s + B] = var
        # 逐维 KL:先对 batch 求和,最后统一除以总样本数
        kl_dim_all += 0.5 * np.sum(mu ** 2 + var - lv - 1.0, axis=0) / n
        # 用同一个 mu 重复采样 n_repeat 次,估计解码器实际看到的抖动有多大
        eps = rng.normal(size=(B, n_repeat, d))
        z = mu[:, None, :] + sig[:, None, :] * eps
        z_flat = z.reshape(-1, d)
        h1d = np.maximum(z_flat @ model.p["V1"] + model.p["c1"], 0.0)
        xhat = (h1d @ model.p["V2"] + model.p["c2"]).reshape(B, n_repeat, D_OBS)
        recon_all += np.sum((xhat - xb[:, None, :]) ** 2) / (n * n_repeat * D_OBS)

    std_mu = mu_all.std(axis=0)                                   # [d]
    mean_sig = np.sqrt(var_all.mean(axis=0))                      # [d]
    marginal_std = np.sqrt(std_mu ** 2 + var_all.mean(axis=0))    # [d]
    overall_std = float(np.sqrt(np.mean(marginal_std ** 2)))
    return dict(
        beta=beta,
        weight_decay=weight_decay,
        recon=float(recon_all),
        kl=float(np.sum(kl_dim_all)),
        kl_dim=kl_dim_all.tolist(),
        std_mu=std_mu.tolist(),
        mean_sig=mean_sig.tolist(),
        marginal_std=marginal_std.tolist(),
        overall_std=overall_std,
        scaling_factor=1.0 / overall_std,
        active_dims=int(np.sum(kl_dim_all > 0.05)),   # 逐维 KL > 0.05 nat 视为「在用」
    )


def sweep(wd_list=None, betas=None, verbose=True):
    wd_list = WD_LIST if wd_list is None else wd_list
    betas = BETAS if betas is None else betas
    X = make_data()
    out = {}
    for wd in wd_list:
        rows = []
        for beta in betas:
            model = train(beta, X, weight_decay=wd)
            rows.append(evaluate(model, X, beta, weight_decay=wd))
        out[f"wd={wd:g}"] = rows
        if verbose:
            print(f"--- weight_decay = {wd:g} ---")
            for m in rows:
                print(f"  beta={m['beta']:<8g} recon={m['recon']:.4f} "
                      f"kl={m['kl']:9.3f}  std(z)={m['overall_std']:7.4f}  "
                      f"1/std={m['scaling_factor']:7.4f}  active={m['active_dims']}/{D_LAT}")
    return out


CACHE = os.path.join(os.path.dirname(os.path.abspath(__file__)), "beta_sweep.json")


def load_or_sweep(path=CACHE):
    """图与正文共用同一份结果,避免图文数字漂移。"""
    if os.path.exists(path):
        with open(path) as f:
            return json.load(f)
    res = sweep()
    with open(path, "w") as f:
        json.dump(res, f, indent=2)
    return res


def _parse_cli(argv):
    """正文 09 节的两个实验就是靠这两个参数复现的:

        python vae_minimal.py --betas 0.03,0.05,0.1,0.2,0.5 --wd 0
        python vae_minimal.py --wd 0.01
    """
    wd_list, betas, full = None, None, False
    for a in argv:
        if a.startswith("--betas="):
            betas = [float(t) for t in a.split("=", 1)[1].split(",")]
        elif a.startswith("--wd="):
            wd_list = [float(t) for t in a.split("=", 1)[1].split(",")]
        elif a == "--full":
            full = True
    return wd_list, betas, full


if __name__ == "__main__":
    import sys

    wd_list, betas, force = _parse_cli(sys.argv[1:])
    custom = wd_list is not None or betas is not None
    if os.path.exists(CACHE) and not force and not custom:
        # 扫描要跑 9 x 2 x 250 = 4500 个 epoch,读缓存是为了让「跑一遍就有输出」
        # 这句话成立;想从头算就加 --full。指定了自定义参数就不走缓存。
        with open(CACHE) as f:
            res = json.load(f)
        print(f"[缓存] 读 {os.path.basename(CACHE)};加 --full 可重跑约 5 分钟的完整扫描")
    else:
        if custom:
            print(f"自定义扫描:wd={wd_list or WD_LIST} beta={betas or BETAS}")
        else:
            print("开始完整扫描(9 个 beta x 2 组权重衰减 x 250 epoch,约 5 分钟)...")
        res = sweep(wd_list, betas)
        if not custom:
            with open(CACHE, "w") as f:
                json.dump(res, f, indent=2)

    print(f"\n{'='*84}")
    print(f"{'weight_decay':>14} {'beta':>9} {'重建MSE':>10} {'总KL':>10} "
          f"{'std(z)':>9} {'1/std':>8} {'存活维度':>8}")
    print(f"{'-'*84}")
    for key, rows in res.items():
        for m in rows:
            print(f"{key:>14} {m['beta']:>9g} {m['recon']:>10.4f} {m['kl']:>10.3f} "
                  f"{m['overall_std']:>9.4f} {m['scaling_factor']:>8.4f} "
                  f"{m['active_dims']:>5}/{D_LAT}")

    print("\n潜尺度拆账(std_mu = 跨维均值变化标准差 RMS;mean_sig = 跨维采样噪声标准差 RMS):")
    for key, rows in res.items():
        print(f"  {key}")
        for m in rows:
            frac = np.mean(np.array(m["std_mu"]) ** 2) / (
                np.mean(np.array(m["std_mu"]) ** 2) + np.mean(np.array(m["mean_sig"]) ** 2))
            print(f"    beta={m['beta']:<8g} std_mu={np.sqrt(np.mean(np.array(m['std_mu']) ** 2)):6.3f} "
                  f"mean_sig={np.sqrt(np.mean(np.array(m['mean_sig']) ** 2)):6.3f} std(z)={m['overall_std']:6.3f} "
                  f"信息占比={frac:5.1%}")

collapse_closed_form.py

"""一维线性 VAE 的闭式解:后验坍缩的阈值到底在哪。

第 03 节推了这么一个结论:

    设  数据 x ~ N(0, v),编码器 q(z|x) = N(m x, s^2),
        解码器 p(x|z) = N(w z, sigma_x^2)(sigma_x 固定,等价于重建项的权重),
        目标 J = [v(1-wm)^2 + w^2 s^2] / (2 sigma_x^2)
                + beta * 0.5 * (m^2 v + s^2 - log s^2 - 1)

    令 lambda = beta * sigma_x^2 / v,则驻点为
        u := w m = 1 - lambda      (潜变量被真正使用的程度)
        s^2     = lambda           (后验方差)
        w^2     = v (1 - lambda)   (解码器权重)
    当 lambda > 1 时此分支要求 w^2 < 0,不存在实数解;lambda = 1 接到坍缩点,最优解退化为坍缩点 (w, m, s) = (0, 0, 1)。

这个脚本用数值优化去拟合 (w, m, s),逐项对照闭式解,
顺便验证「坍缩点永远是一个驻点」这件事。

运行:  /usr/local/bin/python3 collapse_closed_form.py
依赖:  numpy
"""

import numpy as np

V = 1.0            # 数据方差
SIGMA_X = 1.0      # 解码器的观测噪声标准差,固定(它就是重建项的隐式权重)


# ─────────────────────────── 目标函数与梯度 ───────────────────────────
def objective(w, m, s, beta, v=V, sigma_x=SIGMA_X):
    recon = (v * (1.0 - w * m) ** 2 + w ** 2 * s ** 2) / (2.0 * sigma_x ** 2)
    kl = 0.5 * beta * (m ** 2 * v + s ** 2 - np.log(s ** 2) - 1.0)
    return recon + kl


def grad(w, m, s, beta, v=V, sigma_x=SIGMA_X):
    dJ_dw = (-v * m * (1.0 - w * m) + w * s ** 2) / sigma_x ** 2
    dJ_dm = -v * w * (1.0 - w * m) / sigma_x ** 2 + beta * m * v
    dJ_ds = w ** 2 * s / sigma_x ** 2 + beta * (s - 1.0 / s)
    return dJ_dw, dJ_dm, dJ_ds


def fit(beta, v=V, sigma_x=SIGMA_X, steps=20000, lr=0.02, seed=0):
    """对 (w, m, log s) 做梯度下降。用 log s 参数化保证 s > 0。"""
    rng = np.random.default_rng(seed)
    w = rng.normal() * 0.5
    m = rng.normal() * 0.5
    t = rng.normal() * 0.5          # s = exp(t)
    for i in range(steps):
        s = np.exp(t)
        gw, gm, gs = grad(w, m, s, beta, v, sigma_x)
        # 对 t 的梯度要过一次链式法则
        w -= lr * gw
        m -= lr * gm
        t -= lr * gs * s
    return w, m, np.exp(t)


def closed_form(beta, v=V, sigma_x=SIGMA_X):
    if beta <= 0 or v <= 0 or sigma_x <= 0:
        raise ValueError("beta, v and sigma_x must be positive")
    lam = beta * sigma_x ** 2 / v
    if lam >= 1.0:                  # 坍缩
        return dict(lam=lam, u=0.0, s2=1.0, w2=0.0, collapsed=True)
    u = 1.0 - lam
    return dict(lam=lam, u=u, s2=lam, w2=v * u, collapsed=False)


# ─────────────────────────── 主流程 ───────────────────────────
if __name__ == "__main__":
    betas = [0.05, 0.1, 0.2, 0.4, 0.6, 0.8, 0.95, 1.0, 1.2, 2.0, 5.0]
    print(f"v = {V}, sigma_x = {SIGMA_X}, lambda = beta * sigma_x^2 / v = {SIGMA_X**2/V} * beta")
    print(f"闭式预测:坍缩阈值 beta* = v / sigma_x^2 = {V / SIGMA_X**2:g}\n")
    print(f"{'beta':>6} {'lambda':>8} | {'u=wm 预测':>10} {'拟合':>9} "
          f"| {'s^2 预测':>9} {'拟合':>9} | {'w^2 预测':>9} {'拟合':>9} | 坍缩")

    max_err = 0.0
    for beta in betas:
        cf = closed_form(beta)
        w, m, s = fit(beta)
        u_fit = w * m
        err = max(abs(u_fit - cf["u"]), abs(s ** 2 - cf["s2"]), abs(w ** 2 - cf["w2"]))
        max_err = max(max_err, err) if not cf["collapsed"] else max_err
        print(f"{beta:>6g} {cf['lam']:>8.3f} | {cf['u']:>10.4f} {u_fit:>9.4f} "
              f"| {cf['s2']:>9.4f} {s**2:>9.4f} | {cf['w2']:>9.4f} {w**2:>9.4f} "
              f"| {'是' if cf['collapsed'] else '否'}")
    print(f"\n未坍缩区间内,闭式解与数值拟合的最大偏差:{max_err:.2e}")

    # 坍缩点是不是永远的驻点?
    print("\n坍缩点 (w, m, s) = (0, 0, 1) 处的梯度(应当恒为 0):")
    for beta in [0.05, 0.5, 1.0, 5.0]:
        gw, gm, gs = grad(0.0, 0.0, 1.0, beta)
        print(f"  beta={beta:<6g} dJ/dw={gw:+.2e}  dJ/dm={gm:+.2e}  dJ/ds={gs:+.2e}")

    # 坍缩点在不同 beta 下到底是不是最优?比一下目标函数值
    print("\n坍缩点 vs 解析解的目标函数值(谁小谁是最优):")
    for beta in [0.05, 0.5, 0.9, 1.0, 2.0]:
        cf = closed_form(beta)
        if cf["collapsed"]:
            j_star = objective(0.0, 0.0, 1.0, beta)
            j_alt = None
        else:
            w2 = cf["w2"]
            w = np.sqrt(w2)
            m = cf["u"] / w
            j_star = objective(w, m, np.sqrt(cf["s2"]), beta)
            j_alt = objective(0.0, 0.0, 1.0, beta)
        alt = "        —" if j_alt is None else f"{j_alt:9.5f}"
        print(f"  beta={beta:<6g} J(解析解)={j_star:9.5f}  J(坍缩点)={alt}")

latent_scaling.py

"""scaling_factor 到底在补什么:把「潜空间标准差」翻译成「信噪比」。

Stable Diffusion 的 VAE 有个著名常数 0.18215。它是这么用的:

    编码: z_scaled = z * 0.18215        (交给扩散模型之前)
    解码: z        = z_scaled / 0.18215 (交给解码器之前)

diffusers 的文档字符串(AutoencoderKL)写得很明白:这个数是「在训练集第一
个 batch 上算出来的潜空间逐通道标准差」,用它把潜空间缩到单位方差,出处是
LDM 论文(arXiv:2112.10752)的 4.3.2 与 D.1 节。LDM D.1 的原话是:

    "the signal-to-noise ratio induced by the variance of the latent space
     (i.e. Var(z) / sigma_t^2) significantly affects the results ...
     when training a LDM directly in the latent space of a KL-regularized model,
     this ratio is very high ... Note that the VQ-regularized space has a
     variance close to 1, such that it does not have to be rescaled."

本脚本把这段话变成可算的数。核心结论是两条:

1. 若扩散模型是在单位方差潜空间上训练的,它对 z_0 的后验均值估计就是
        z_hat_0 = sqrt(alpha_bar) * z_t
   (高斯先验 N(0, I) 下的 MMSE 估计)。而真实潜空间标准差是 sigma 时,
   正确的 MMSE 估计是
        z_hat_0 = sigma^2 sqrt(alpha_bar) / (sigma^2 alpha_bar + 1 - alpha_bar) * z_t
   两者之比 r = alpha_bar + (1 - alpha_bar) / sigma^2。
   r < 1 就是「重建出来的东西被整体缩小」——图发灰。

2. 每一层的真实信噪比是 sigma^2 * alpha_bar / (1 - alpha_bar),
   模型以为的是 alpha_bar / (1 - alpha_bar),整整差 sigma^2 倍
   (20*log10(sigma) 分贝)。

运行:  /usr/local/bin/python3 latent_scaling.py
依赖:  numpy
"""

import numpy as np

SD_SCALING_FACTOR = 0.18215        # diffusers AutoencoderKL 的默认值
ABARS = [0.999, 0.99, 0.9, 0.5, 0.1, 0.01]


def latent_std_from_scaling(s):
    """scaling_factor 是潜空间边际标准差的倒数。"""
    return 1.0 / s


def mmse_model(z_t, abar):
    """在「潜空间方差 = 1」的假设下训练出来的高斯 MMSE 去噪器。"""
    return np.sqrt(abar) * z_t


def mmse_true(z_t, abar, sigma):
    """潜空间真实标准差为 sigma 时的高斯 MMSE 去噪器。"""
    return sigma ** 2 * np.sqrt(abar) / (sigma ** 2 * abar + 1.0 - abar) * z_t


def amplitude_ratio(abar, sigma):
    """模型输出 / 正确输出 = alpha_bar + (1 - alpha_bar) / sigma^2。"""
    return abar + (1.0 - abar) / sigma ** 2


def snr(abar, sigma):
    """真实信噪比 Var(z_0 的信号成分) / Var(噪声成分) = sigma^2 * abar / (1 - abar)。"""
    return sigma ** 2 * abar / (1.0 - abar)


def mmse_mse(abar, sigma):
    """高斯 MMSE 估计的理论误差,就是后验方差。"""
    return sigma ** 2 * (1.0 - abar) / (sigma ** 2 * abar + 1.0 - abar)


def demo_step(abar, sigma, n=400000, seed=0):
    """蒙特卡洛:固定同一个前向过程,比较两个去噪器。

    est_model:按「潜空间方差 = 1」训练出来的理想去噪器,用在方差为 sigma^2 的
               潜空间上 —— 这就是拿了别人训好的权重却没乘 scaling_factor 的情形。
    est_true :就该 sigma 训练出来的理想去噪器 —— 也就是正确缩放后应有的表现。
    """
    rng = np.random.default_rng(seed)
    z0 = sigma * rng.normal(size=n)
    eps = rng.normal(size=n)
    z_t = np.sqrt(abar) * z0 + np.sqrt(1.0 - abar) * eps
    est_model = mmse_model(z_t, abar)
    est_true = mmse_true(z_t, abar, sigma)
    return (float(np.mean((est_model - z0) ** 2)),
            float(np.mean((est_true - z0) ** 2)))


if __name__ == "__main__":
    import sys

    # 正文 09 节实验三:换一个反事实的 scaling_factor 重算整张表
    #   python latent_scaling.py --sf=0.5
    #   python latent_scaling.py --sf=2.0
    for a in sys.argv[1:]:
        if a.startswith("--sf="):
            SD_SCALING_FACTOR = float(a.split("=", 1)[1])
    sigma_sd = latent_std_from_scaling(SD_SCALING_FACTOR)
    print("=" * 72)
    print("一、0.18215 意味着什么")
    print("=" * 72)
    print(f"  scaling_factor = {SD_SCALING_FACTOR}")
    print(f"  潜空间边际标准差 sigma = 1 / {SD_SCALING_FACTOR} = {sigma_sd:.4f}")
    print(f"  潜空间边际方差 Var(z)  = {sigma_sd**2:.4f}")
    print(f"  信噪比偏移             = {20*np.log10(sigma_sd):.2f} dB")
    print()
    print("  对照:若某个 VAE 的潜空间方差真的接近 1(LDM 说 VQ 版就是如此),")
    print("  scaling_factor 就应当接近 1,完全不需要 rescale。")

    print()
    print("=" * 72)
    print("二、忘掉 scaling factor,去噪幅度错多少")
    print("=" * 72)
    print("  r = 模型输出 / 正确输出;r=1 才是正确的")
    print(f"  {'alpha_bar':>10} | " + " | ".join(f"sigma={s:<7.4g}" for s in [sigma_sd, 1.0, SD_SCALING_FACTOR]))
    for abar in ABARS:
        rs = [amplitude_ratio(abar, s) for s in [sigma_sd, 1.0, SD_SCALING_FACTOR]]
        print(f"  {abar:>10g} | " + " | ".join(f"{r:>13.4f}" for r in rs))

    print()
    print("=" * 72)
    print("三、蒙特卡洛实测:单步去噪的均方误差(n=400000)")
    print("=" * 72)
    print(f"  {'alpha_bar':>10} | {'错假设 MSE':>12} {'正确 MSE':>12} {'理论后验方差':>14} {'恶化倍数':>10}")
    for abar in ABARS:
        m_bad, m_good = demo_step(abar, sigma_sd)
        print(f"  {abar:>10g} | {m_bad:>12.4f} {m_good:>12.4f} {mmse_mse(abar, sigma_sd):>14.4f} "
              f"{m_bad/m_good:>10.2f}x")
    print()
    print("  读法:固定同一个前向过程 z_t = sqrt(alpha_bar) z_0 + sqrt(1-alpha_bar) eps,")
    print("  「错假设」是拿了按单位方差潜空间训好的去噪器却喂未缩放的 latent,")
    print("  「正确」是就该 sigma 训练的去噪器(也就是乖乖乘上 0.18215 的效果)。")
    print(f"  「理论后验方差」= sigma^2 (1-alpha_bar) / (sigma^2 alpha_bar + 1 - alpha_bar),")
    print(f"  与蒙特卡洛的「正确」一列应当吻合(校核 alpha_bar=0.5:"
          f"{mmse_mse(0.5, sigma_sd):.4f} vs 实测 {demo_step(0.5, sigma_sd)[1]:.4f})")

    print()
    print("=" * 72)
    print("四、反向的错误:换了个方差接近 1 的 VAE,却照抄 0.18215")
    print("=" * 72)
    print(f"  {'alpha_bar':>10} | {'r(幅度被放大的倍数)':>24}")
    for abar in ABARS:
        print(f"  {abar:>10g} | {amplitude_ratio(abar, SD_SCALING_FACTOR):>24.4f}")
    print()
    print("  r >> 1:重建幅度被整体放大,对应生成图过曝、结构糊成一团。")

    print()
    print("=" * 72)
    print("五、信噪比视角(LDM D.1 的定义 Var(z) / sigma_t^2)")
    print("=" * 72)
    print(f"  {'alpha_bar':>10} | {'模型以为的 SNR':>16} {'真实 SNR':>16} {'倍数':>10}")
    for abar in ABARS:
        a = abar / (1.0 - abar)
        b = snr(abar, sigma_sd)
        print(f"  {abar:>10g} | {a:>16.4f} {b:>16.4f} {b/a:>10.2f}x")

make_figures.py

"""画本文的四张图。数据源全部来自 vae_minimal.py / latent_scaling.py 的真实输出,
不另造数,避免图文数字漂移。

    beta_sweep.png          beta 扫描全景:重建 / KL / 潜尺度 / 存活维度
    latent_scale_split.png  潜尺度的拆账:编码位置 vs 采样噪声
    scaling_snr.png         忘掉 scaling_factor 的后果:幅度比 + 单步去噪 MSE
    collapse_profile.png    后验坍缩的指纹:逐维 KL 怎么一个个死掉

运行:  /usr/local/bin/python3 make_figures.py
依赖:  numpy, matplotlib(字体 PingFang SC)
"""

import os
import numpy as np
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt

import vae_minimal as VM
import latent_scaling as LS

plt.rcParams["font.sans-serif"] = ["PingFang SC", "Arial Unicode MS", "Heiti TC", "sans-serif"]
plt.rcParams["axes.unicode_minus"] = False
plt.rcParams["figure.facecolor"] = "white"
plt.rcParams["axes.facecolor"] = "white"
plt.rcParams["savefig.facecolor"] = "white"
plt.rcParams["font.size"] = 11

HERE = os.path.dirname(os.path.abspath(__file__))
FIGDIR = os.path.join(os.path.dirname(HERE), "figures")
os.makedirs(FIGDIR, exist_ok=True)          # exist_ok 必须带,否则目录已存在会 PermissionError

C_MAIN, C_ALT = "#2563eb", "#dc2626"
C_GRAY = "#6b7280"


def _b(rows):
    return np.array([r["beta"] for r in rows])


# ─────────────────────────── 图 1:beta 扫描全景 ───────────────────────────
def fig_beta_sweep(res):
    rows0 = res["wd=0"]
    rows1 = res["wd=0.001"]
    fig, axes = plt.subplots(2, 2, figsize=(13.5, 8.6))
    ax = axes[0][0]
    ax.plot(_b(rows0), [r["recon"] for r in rows0], "o-", color=C_MAIN, label="无权重衰减")
    ax.plot(_b(rows1), [r["recon"] for r in rows1], "s--", color=C_ALT, label="权重衰减 1e-3")
    ax.set_xscale("log")
    ax.set_yscale("log")
    ax.axhline(1.0, color=C_GRAY, lw=1, ls=":")
    ax.text(1.1e-6, 1.05, "重建 MSE = 1:等于什么都没学到", color=C_GRAY, fontsize=9)
    ax.set_xlabel(r"KL 权重 $\beta$")
    ax.set_ylabel("重建 MSE($D$ 维平均)")
    ax.set_title("(a) 重建质量:β 一大就断崖")
    ax.legend(fontsize=9)
    ax.grid(alpha=0.25)

    ax = axes[0][1]
    ax.plot(_b(rows0), [r["kl"] for r in rows0], "o-", color=C_MAIN, label="无权重衰减")
    ax.plot(_b(rows1), [r["kl"] for r in rows1], "s--", color=C_ALT, label="权重衰减 1e-3")
    ax.set_xscale("log")
    ax.set_yscale("symlog", linthresh=1e-1)
    ax.set_xlabel(r"KL 权重 $\beta$")
    ax.set_ylabel("总 KL(nats,8 维求和)")
    ax.set_title("(b) KL:β 越大越贴先验,坍缩时归零")
    ax.legend(fontsize=9)
    ax.grid(alpha=0.25)

    ax = axes[1][0]
    ax.plot(_b(rows0), [r["overall_std"] for r in rows0], "o-", color=C_MAIN, label="无权重衰减")
    ax.plot(_b(rows1), [r["overall_std"] for r in rows1], "s--", color=C_ALT, label="权重衰减 1e-3")
    ax.axhline(1.0, color=C_GRAY, lw=1, ls=":")
    ax.text(1.1e-6, 1.03, "std(z) = 1:被 KL 钉住", color=C_GRAY, fontsize=9)
    ax.axhline(LS.latent_std_from_scaling(LS.SD_SCALING_FACTOR), color="#059669", lw=1, ls="--")
    ax.text(1.1e-6, 5.62, "std(z) = 5.49:由 SD1.x 配置系数反推", color="#059669", fontsize=9)
    ax.set_xscale("log")
    ax.set_xlabel(r"KL 权重 $\beta$")
    ax.set_ylabel("潜空间边际标准差 std(z)")
    ax.set_ylim(0.9, 6.2)
    ax.set_title("(c) 潜尺度:β 一松手就漂走")
    ax.legend(fontsize=9, loc="upper right")
    ax.grid(alpha=0.25)

    ax2 = ax.twinx()
    ax2.set_ylim(1 / 6.2, 1 / 0.9)
    ax2.set_ylabel("对应的 scaling_factor = 1 / std(z)")

    ax = axes[1][1]
    ax.plot(_b(rows0), [r["active_dims"] for r in rows0], "o-", color=C_MAIN, label="无权重衰减")
    ax.plot(_b(rows1), [r["active_dims"] for r in rows1], "s--", color=C_ALT, label="权重衰减 1e-3")
    ax.axhline(4, color="#059669", lw=1, ls="--")
    ax.text(1.1e-6, 4.2, "真实因子数 = 4", color="#059669", fontsize=9)
    ax.set_xscale("log")
    ax.set_xlabel(r"KL 权重 $\beta$")
    ax.set_ylabel("还活着的维度数(逐维 KL > 0.05 nat)")
    ax.set_ylim(-0.3, 8.6)
    ax.set_title("(d) 维度预算:冗余维度被逐个关掉")
    ax.legend(fontsize=9, loc="lower left")
    ax.grid(alpha=0.25)

    fig.suptitle("KL 权重扫描:重建、KL、潜尺度、存活维度(潜变量 8 维,真实因子 4 个)",
                 fontsize=13, y=0.98)
    fig.tight_layout(rect=[0, 0, 1, 0.96])
    out = os.path.join(FIGDIR, "beta_sweep.png")
    fig.savefig(out, dpi=150)
    plt.close(fig)
    return out


# ─────────────────────────── 图 2:潜尺度拆账 ───────────────────────────
def fig_scale_split(res):
    rows = res["wd=0"]
    betas = _b(rows)
    std_mu = np.array([np.mean(np.array(r["std_mu"]) ** 2) for r in rows])        # 编码位置的跨样本波动
    mean_sig = np.array([np.mean(np.array(r["mean_sig"]) ** 2) for r in rows])    # 采样噪声 RMS
    x = np.arange(len(betas))

    fig, axes = plt.subplots(1, 2, figsize=(13.5, 4.6))
    ax = axes[0]
    ax.bar(x, std_mu, 0.62, label=r"后验均值的跨样本方差", color=C_MAIN)
    ax.bar(x, mean_sig, 0.62, bottom=std_mu,
           label=r"平均后验方差", color="#f59e0b")
    ax.set_xticks(x)
    ax.set_xticklabels([f"{b:g}" for b in betas])
    ax.set_xlabel(r"KL 权重 $\beta$(对数轴上的等距刻度)")
    ax.set_ylabel("Var(z) 的构成(各维平均)")
    ax.set_title("(a) 方差拆账:后验均值变化与后验噪声")
    ax.legend(fontsize=9)
    ax.grid(alpha=0.25, axis="y")

    ax = axes[1]
    frac = std_mu / (std_mu + mean_sig)
    ax.plot(x, frac, "o-", color=C_MAIN, lw=2)
    ax.axhline(0.5, color=C_GRAY, lw=1, ls=":")
    ax.set_xticks(x)
    ax.set_xticklabels([f"{b:g}" for b in betas])
    ax.set_xlabel(r"KL 权重 $\beta$")
    ax.set_ylabel("方差里来自编码位置的比例")
    ax.set_ylim(-0.03, 1.06)
    ax.set_title("(b) 总方差中均值变化的占比")
    ax.grid(alpha=0.25)
    for i, (xi, fi) in enumerate(zip(x, frac)):
        ax.annotate(f"{fi:.2f}", (xi, fi), textcoords="offset points",
                    xytext=(0, 8), ha="center", fontsize=9, color=C_MAIN)

    fig.suptitle("全方差公式:两项方差相加;方差占比不等于互信息",
                 fontsize=12.5, y=1.0)
    fig.tight_layout(rect=[0, 0, 1, 0.95])
    out = os.path.join(FIGDIR, "latent_scale_split.png")
    fig.savefig(out, dpi=150)
    plt.close(fig)
    return out


# ─────────────────────────── 图 3:scaling factor 的后果 ───────────────────────────
def fig_scaling_snr():
    sigma_sd = LS.latent_std_from_scaling(LS.SD_SCALING_FACTOR)
    abar = np.logspace(-3, np.log10(0.9999), 400)

    fig, axes = plt.subplots(1, 2, figsize=(13.5, 4.8))
    ax = axes[0]
    for sigma, lab, col in [(sigma_sd, rf"$\sigma$={sigma_sd:.2f}(由 SD1.x 配置反推)", C_MAIN),
                            (1.0, r"$\sigma$=1.0(正确缩放后)", "#059669"),
                            (LS.SD_SCALING_FACTOR, rf"$\sigma$={LS.SD_SCALING_FACTOR:.3f}(照抄 0.18215 用错 VAE)", C_ALT)]:
        ax.plot(abar, LS.amplitude_ratio(abar, sigma), "-", color=col, lw=2, label=lab)
    ax.axhline(1.0, color=C_GRAY, lw=1, ls=":")
    ax.text(2e-3, 1.35, "$r=1$:缩放正确", color=C_GRAY, fontsize=9)
    ax.text(2e-3, 0.038, r"$r=1/\sigma^2=0.033$:估计幅度约为正确值的 3.3%", color=C_MAIN, fontsize=9)
    ax.set_xscale("log")
    ax.set_yscale("log")
    ax.set_ylim(2e-2, 6e1)
    ax.set_xlabel(r"$\bar{\alpha}$(1 = 完全干净,0 = 纯噪声)")
    ax.set_ylabel("去噪幅度比 $r$ = 模型输出 / 正确输出")
    ax.set_title("(a) 高斯 MMSE 教学模型的后验均值幅度比")
    ax.legend(fontsize=9, loc="center left")
    ax.grid(alpha=0.25, which="both")

    ax = axes[1]
    abars = np.array(LS.ABARS)
    bad, good = [], []
    for a in abars:
        b_, g_ = LS.demo_step(a, sigma_sd, n=200000)
        bad.append(b_)
        good.append(g_)
    ax.plot(abars, bad, "o-", color=C_ALT, lw=2, label=r"错假设去噪器(按 $\sigma$=1 训练)")
    ax.plot(abars, good, "s-", color=C_MAIN, lw=2, label=r"正确去噪器(就该 $\sigma$ 训练)")
    ax.set_xscale("log")
    ax.set_yscale("log")
    ax.set_xlabel(r"$\bar{\alpha}$")
    ax.set_ylabel("单步去噪均方误差")
    ax.set_title("(b) 同一个前向过程下的 MSE 差距")
    ax.legend(fontsize=9)
    ax.grid(alpha=0.25)
    i5 = int(np.argmin(np.abs(abars - 0.5)))
    ax.annotate(f"{bad[i5]/good[i5]:.1f}×", (abars[i5], bad[i5]),
                textcoords="offset points", xytext=(12, -6), color=C_ALT, fontsize=11)

    fig.suptitle(r"忘掉 scaling_factor 的代价:信噪比被高估 $\sigma^2$ = 30.1 倍(14.8 dB)",
                 fontsize=12.5, y=1.0)
    fig.tight_layout(rect=[0, 0, 1, 0.94])
    out = os.path.join(FIGDIR, "scaling_snr.png")
    fig.savefig(out, dpi=150)
    plt.close(fig)
    return out


# ─────────────────────────── 图 4:坍缩指纹 ───────────────────────────
def fig_collapse_profile(res):
    rows = res["wd=0"]
    picks = [3e-3, 3e-1, 1.0]     # 三档:够用 / 正在坍缩 / 完全坍缩
    # 每张子图各自缩放:三档之间差 20 倍,共享 y 轴会把另两张压成一条线
    fig, axes = plt.subplots(1, 3, figsize=(13.5, 4.2))
    d = len(rows[0]["kl_dim"])
    for ax, beta in zip(axes, picks):
        row = min(rows, key=lambda r: abs(r["beta"] - beta))
        kl = np.array(row["kl_dim"])
        colors = [C_MAIN if v > 0.05 else "#d1d5db" for v in kl]
        ax.bar(np.arange(d), kl, 0.68, color=colors)
        ax.axhline(0.05, color=C_GRAY, lw=1, ls=":")
        ax.set_xticks(np.arange(d))
        ax.set_xticklabels([f"$z_{j+1}$" for j in range(d)], fontsize=9)
        ax.set_xlabel("潜变量维度")
        ax.set_ylim(0, max(0.35, kl.max() * 1.18))
        ax.set_title(rf"$\beta$={row['beta']:g} 存活 {row['active_dims']}/{d} "
                     f"重建 MSE={row['recon']:.4f}")
        ax.grid(alpha=0.25, axis="y")
        ax.text(0.99, 0.92, f"纵轴最大 {kl.max():.2f} nats", transform=ax.transAxes,
                ha="right", fontsize=8.5, color=C_GRAY)
    axes[0].set_ylabel("逐维 KL(nats)")
    axes[1].set_ylabel("逐维 KL(nats)")
    axes[2].set_ylabel("逐维 KL(nats)")
    fig.suptitle("后验坍缩的指纹:KL 预算被削减时,维度是一个个死的,不是一起死(注意三张图的纵轴量级不同)",
                 fontsize=12.5, y=0.99)
    fig.tight_layout(rect=[0, 0, 1, 0.91])
    out = os.path.join(FIGDIR, "collapse_profile.png")
    fig.savefig(out, dpi=150)
    plt.close(fig)
    return out


if __name__ == "__main__":
    res = VM.load_or_sweep()
    for fn in [fig_beta_sweep, fig_scale_split, fig_collapse_profile]:
        print("写出", fn(res))
    print("写出", fig_scaling_snr())
0

评论 (0)

取消
粤ICP备2021042327号