首页
应用
关于
Search
1
Pytorch DDP
2,482 阅读
2
Pytorch 常见问题
1,515 阅读
3
视频时序切分
1,341 阅读
4
中文场景下的CLIP图文预训练
1,044 阅读
5
Semi-Supervised + Noisy Label
1,028 阅读
AIGC
AIGC Daily Papers
AIGC Fundamentals
其他
职场经验复盘
多模态理解
购房/投资
阅读
广告
Segmentation
LeetCode
Pytorch
Python
Shell
C++
常用链接
Search
标签搜索
AIGC
人工智能
论文速读
ai
视频生成
DiT
对齐
蒸馏
扩散模型
attention
transformer
图像生成
视频编辑
diffusion
基础知识
稀疏注意力
多模态
文生图
NVIDIA
llm
Jefxiong
累计撰写
205
篇文章
累计收到
8
条评论
首页
应用
栏目
AIGC
AIGC Daily Papers
AIGC Fundamentals
其他
职场经验复盘
多模态理解
购房/投资
阅读
广告
Segmentation
LeetCode
Pytorch
Python
Shell
C++
常用链接
页面
关于
搜索到
34
篇与
AIGC Fundamentals
的结果
2026-09-28
AIGC 基本功|流匹配与 Rectified Flow-FlowMatching
流匹配与 Rectified Flow 所属方向:生成范式 | 难度:进阶 | 前置知识:DDPM 的训练目标与采样流程(ddpm)、扩散过程的前向与反向推导(diffusion_math) 关键词:流匹配、flow matching、rectified flow、速度场、直线路径、reflow、NFE 01. 为什么需要它 先给三个数字,全部来自文末附录里能直接跑的脚本。 数字一:换个回归目标,预测误差幅度系数差 100 倍,对应平方损失权重差 10000 倍。 把「干净图估计」上的误差记为 $\delta$,在 $t=0.99$(几乎纯噪声)这一时刻:如果模型输出的是噪声 $\epsilon$,loss 看到的误差是 $0.0101\,\delta$;如果输出的是速度 $v$,loss 看到的是 $1.0101\,\delta$。在固定干净图误差的比较下,速度目标给予高噪声端更大的相对权重。 这是速度参数化与噪声参数化的一个权重差异;采用何种目标还取决于路径、预条件、训练时间分布与架构,不能把所有模型选择归因为这一个系数(03 节推导,图 1 画出整条曲线)。 数字二:「直线路径」的 ODE 轨迹一点都不直。 Rectified Flow 最常被转述成"路径是直的,所以快"。实测:把二维八模高斯混合推到数据端,用附录 reflow_lab.py 的 10 万个起点、64 步 RK2(128 NFE)轨迹账本,精确场的弧长/弦长均值为 1.6499(图 2 另用 300 步 RK4 示意)——轨迹比两点连线多走了 65% 的路。训练时那条 $(1-t)x+t\epsilon$ 插值线确实是直线,但模型实际采样走的是边缘速度场的积分曲线,两者不是一回事(02 节的图 2 把这件事画出来了)。 数字三:reflow 一轮,2 步采样的误差降 11.1 倍。 用第 1 轮训好的模型把噪声推到数据端、拿得到的配对重训一轮,配对插值上的归一化回归残差 $S_{pair}$ 从 $0.6229$ 掉到 $0.0027$(约 231 倍;它不是几何曲率),NFE=2 的终点偏差从 $0.04469$ 降到 $0.00402$,已经低于方差保持路径 2 步的 $0.0105$。但高 NFE 下的终点偏差没有因此继续下降:生成质量仍受教师分布、重训误差和有限样本评估影响,06 节展开。 所以这篇文章要回答三个问题:流匹配到底在训练什么、它和 DDPM 是不是两个东西、以及"直线路径"这个卖点真实兑现了多少。 02. 最小可用理解 三句话讲完: 流匹配训练的是一个速度场,不是噪声。 采样是解一条 ODE:从噪声端 $t=1$ 出发,跟着 $v_\theta$ 走到数据端 $t=0$。训练只是回归:给定当前点 $z_t$,预测条件速度 $u_t$;在线性 Rectified Flow 路径上,它就是这一对端点的相对位移。最优解是条件期望 $E[u_t \mid z_t = z]$,而由连续性方程(03 节证),这个期望场恰好把 $p_1$ 运到 $p_0$——所以不需要知道任何密度,样本对就够。 高斯路径提供了统一描述,但不同路径不等于同一个模型。 把路径统一写成 $z_t=\alpha_t x+\sigma_t\epsilon$:DDPM 那一路取 $\alpha_t=\sqrt{\bar\alpha_t}$、$\sigma_t=\sqrt{1-\bar\alpha_t}$;Rectified Flow 取 $\alpha_t=1-t$、$\sigma_t=t$。两者可以使用相同的网络架构。三种预测目标(干净图 $x$、噪声 $\epsilon$、速度 $v$)在同一条已知路径及非退化时间内可以代数换算输出(无需重训网络);不同路径的模型不能仅换公式就变成彼此,04 节把换算残差验到 $10^{-15}$。 "直"的是训练插值线,不是采样轨迹。 在本文独立配对的平滑数据分布上,线性路径边缘场满足 $v(z,0)=-z$、$v(z,1)=z-\mu$,其中 $\mu=E[x]$。采样从 1 积分到 0,时间步为负,因此噪声端先朝数据均值走,数据端局部向外走;中间还有模式分流。实际积分轨迹可以弯曲,reflow 用模型生成的端点配对重训,是减小这种弯曲的一种办法。 这张图要看什么:在这 14 个共同起点上,左图(直线路径)出现明显回转,右图(VP 余弦路径)较平缓。灰点是 8 个数据模式,蓝圆是噪声起点,绿方是生成终点。两图使用相同的精确场计算方式和 300 步 RK4,只改变路径;其形状是本分布上的实验结果。VP 数据端速度为零,但由于数据均值不为零,噪声端速度并不为零,见 3.4 节。 03. 数学推导 3.1 高斯路径与条件速度 把常见高斯插值类扩散/流方法统一成一条路径。取数据样本 $x\sim p_{\text{data}}$、独立噪声 $\epsilon\sim N(0,I)$,令 $$z_t=\alpha_t x+\sigma_t\epsilon,\qquad t\in[0,1]$$ 约定 $t=0$ 是数据端、$t=1$ 是噪声端,即 $\alpha_0=1,\sigma_0=0$、$\alpha_1=0,\sigma_1=1$。各符号的含义:$\alpha_t$ 是数据分量的幅度,$\sigma_t$ 是噪声分量的幅度,两者是标量函数,选不同的曲线就得到不同的方法: 方法 $\alpha_t$ $\sigma_t$ 备注 VP / DDPM $\sqrt{\bar\alpha_t}$ $\sqrt{1-\bar\alpha_t}$ 满足 $\alpha_t^2+\sigma_t^2=1$,当数据协方差也为单位阵时保持单位方差;一般数据协方差仍随 t 改变 Rectified Flow $1-t$ $t$ 线性插值,中间分布方差会缩水(06 节) 对固定的一对 $(x,\epsilon)$,$z_t$ 是一条确定的曲线,它对时间的导数是 $$u_t=\dot\alpha_t x+\dot\sigma_t\epsilon$$ 每个符号:$\dot\alpha_t$、$\dot\sigma_t$ 是两条幅度曲线的导数;$u_t$ 叫条件速度——它说的是"这一对端点对应的粒子此刻在往哪走"。直线路径下 $\dot\alpha_t=-1$、$\dot\sigma_t=1$,所以 $u_t=\epsilon-x$,与 $t$ 无关:整条插值线是匀速直线。这就是"直线"的全部含义,它只说了这条以 $(x,\epsilon)$ 为参数的曲线,没说采样时走的那条。 3.2 边缘速度场:为什么样本对就够训练 采样时我们手里没有那对 $(x,\epsilon)$,只有 $z_t$。所以要用的是边缘速度场: $$v^*(z,t)=E\left[u_t \mid z_t=z\right]$$ 读法:在所有"此刻恰好经过 $z$ 的粒子对"里,平均下来往哪走。要证它是"对的场",也就是用它解 ODE 确实把 $p_1$ 运到 $p_0$。对任意光滑测试函数 $\phi$,沿条件路径求导再取期望: $$\frac{d}{dt}E\left[\phi(z_t)\right]=E\left[\nabla\phi(z_t)\cdot u_t\right]=E\left[\nabla\phi(z_t)\cdot E[u_t\mid z_t]\right]=\int \nabla\phi(z)\cdot v^*(z,t)\,p_t(z)\,dz$$ 第二步是把里层的条件期望提出来(塔性质)。最后那个积分正是连续性方程 $\partial_t p_t+\nabla\cdot(v^*p_t)=0$ 的弱形式,在速度场和密度满足适当正则性、相应 ODE 与连续性方程有唯一解的条件下,$v^*$ 生成的流在每个时刻匹配既定的 $p_t$,从 $p_1$ 走到 $p_0$。这里 $p_t$ 本身随时间变化,并不是保持某个固定分布不变。这一步就是"marginalization trick":条件场逐对可得,边缘场取个条件期望就行。 于是在线性路径下,训练目标可以完全绕开密度(一般路径把目标换成 $u_t$): $$\min_\theta\;E_{t,x,\epsilon}\left\|v_\theta(z_t,t)-(\epsilon-x)\right\|^2$$ 固定 $t$ 后,这个回归的逐点最优解恰是 $v^*(z,t)$。没有任何一项需要 $p_t$ 的表达式——这是流匹配在工程上能起飞的根本原因:它把"学分布"变成了"学一个回归"。 3.3 同一路径的三种预测输出如何换算 设 $m(z,t)=E[x\mid z_t=z]$(干净图估计)、$e(z,t)=E[\epsilon\mid z_t=z]$(噪声估计)。对 $z_t=\alpha_t x+\sigma_t\epsilon$ 两边取条件期望(条件期望是线性的,$z_t$ 在条件下是常数): $$z=\alpha_t m+\sigma_t e$$ 这是贯穿全文的恒等式,$\alpha m+\sigma e=z$。再配合速度的定义 $v=\dot\alpha_t m+\dot\sigma_t e$,两个方程、两个未知数,解出 $$m=\frac{\sigma_t v-\dot\sigma_t z}{\dot\alpha_t\sigma_t-\dot\sigma_t\alpha_t},\qquad e=\frac{\dot\alpha_t z-\alpha_t v}{\dot\alpha_t\sigma_t-\dot\sigma_t\alpha_t}$$ 分母 $\dot\alpha_t\sigma_t-\dot\sigma_t\alpha_t$ 是个行列式。直线路径下它恒等于 $-1$($(-1)\cdot t-1\cdot(1-t)=-1$),于是换算干净得没有任何除法: 已知 求 $m$(干净图) 求 $e$(噪声) 速度 $v$ $m=z-t\,v$ $e=z+(1-t)\,v$ 噪声 $\epsilon$ $m=(z-t\,e)/(1-t)$ — 干净图 $m$ — $e=(z-(1-t)m)/t$ 第一行值得盯一眼:从速度换算到干净图和噪声都不需要除以 $\alpha$ 或 $\sigma$,而第二、三行各有一个会趋奇的除法。这就是速度参数化在数值上更稳的那一半理由。 那为什么不能说"三种参数化完全等价"? 真值层面等价(上表),loss 层面不等价。假设模型在 $m$ 上有误差 $\delta$,把它代入三种目标的误差: $$\text{x0 目标:}\ \|\delta\|,\qquad \text{噪声目标:}\ \frac{\alpha_t}{\sigma_t}\|\delta\|,\qquad \text{速度目标:}\ \left|\dot\alpha_t-\frac{\dot\sigma_t\alpha_t}{\sigma_t}\right|\|\delta\|$$ 来历:$e=(z-\alpha_t m)/\sigma_t$ 对 $m$ 求导得 $-\alpha_t/\sigma_t$;$v=\dot\alpha_t m+\dot\sigma_t e$ 对 $m$ 求导得 $\dot\alpha_t-\dot\sigma_t\alpha_t/\sigma_t$。直线路径代入,得到三条曲线 $$\text{x0:}1,\qquad \text{噪声:}\frac{1-t}{t},\qquad \text{速度:}\frac{1}{t}$$ 这张图要看什么:横轴是时间 $t$,纵轴是对数刻度的放大倍数。$t=0.99$(接近纯噪声端)噪声参数化的曲线是 $0.0101$——在高噪声端,噪声目标对"你把干净图猜错了"几乎无感,因为 $z$ 里本来就几乎没有 $x$ 的信息;而速度目标的系数是 $1.0101$;真正取 $t\to1$ 的极限时,二者分别趋于 0 和 1。反过来在 $t\to 0$(数据端)二者对固定干净图误差的换算系数都增大,速度系数比噪声系数加法上大 1;这不等于直接 velocity 训练的目标或梯度必然发散。这里能直接推出的是高噪声端对干净图估计误差的相对损失权重不同,不能单凭它保证真实网络的梯度大小或学习效率。 3.4 直的是插值线,不是轨迹 先看端点行为。本文数据是平滑的高斯混合,数据与噪声独立;因此在数据端,$m(z,0)=z$、$e(z,0)=E[\epsilon\mid x=z]=0$;在噪声端,$m(z,1)=\mu=E[x]$、$e(z,1)=z$。这些关系来自端点的条件独立性,不能用 $(z-\alpha m)/\sigma$ 中的“分子趋零”来推断一个 $0/0$ 极限。 make_data 的 8 个分量权重与 $1+0.35\cos(\cdot)$ 成正比,并不均匀。精确加权均值为 $$\mu=(0.294464,\,-0.248024).$$ 线性路径 $\dot\alpha=-1,\dot\sigma=1$ 的端点速度是 $$v(z,0)=-z,\qquad v(z,1)=z-\mu.$$ 采样沿负时间方向积分,所以噪声端局部朝 $\mu$ 走,数据端局部沿 $z$ 向外走。中间的模式分流也会影响轨迹,端点公式本身不能决定全程弯曲程度;01 节的弧长比是实际积分测得的结果,不是由端点方向推出的普遍定理。 VP 余弦路径取 $\alpha=\cos(\pi t/2)$、$\sigma=\sin(\pi t/2)$,因此 $$v_{\rm VP}(z,0)=0,\qquad v_{\rm VP}(z,1)=-\frac{\pi}{2}\mu=(-0.462543,\,0.389595).$$ 只有数据均值为零时,这条 VP 路径的两个端点速度才都为零。本实验不满足该条件。 需要区分两个诊断量。本文脚本 straightness_of 实际计算的是原始端点配对插值上的归一化回归残差: $$S_{pair}=\frac{\int_0^1 E\|v((1-t)z_0+t z_1,t)-(z_1-z_0)\|^2dt}{E\|z_1-z_0\|^2}.$$ 当 $v=v^*$ 时,分子是条件回归的不可约误差。下式的期望同时包含均匀时间与端点配对,并且 $S_{pair}$ 用 $v^*$ 计算: $$E\|v_\theta-u\|^2=E\|v_\theta-v^*\|^2+S_{pair}E\|z_1-z_0\|^2.$$ Rectified Flow 论文沿实际 ODE 轨迹 $Z_t$ 定义的 straightness 则比较 $v(Z_t,t)$ 与 $Z_1-Z_0$。这不是在独立原始配对插值线上评估,不能直接把本文 $S_{pair}$ 当成论文的同名量。本文另用实际积分轨迹的弧长/弦长检测几何变直,两个指标应分别报告。弧长的离散计算必须是 $\sum_j\|Z_{t_{j+1}}-Z_{t_j}\|_2$,不能先分别累加坐标上的绝对位移再取范数。 path_lab.py 另取 1500 个起点、400 步 RK4(1600 NFE):线性路径的平均弧长/弦长为 1.6232,VP 为 1.2915。沿实际轨迹的速度偏差再除以平均弦长平方,得到归一化 $S_{traj}$ 分别为 1.2293、0.5317。这与原始配对插值上的 $S_{pair}$ 不同。01 节的 1.6499 来自 reflow_lab.py 的另一批 10 万个起点和 64 步 RK2;样本与积分精度不同,诊断数值也会不同。 04. 代码实现 全部代码在文末附录,五个脚本,数值实验依赖 numpy、画图另需 matplotlib:fm_oracle.py(实验台与精确场)、param_lab.py(三种参数化)、path_lab.py(路径对比与 NFE)、reflow_lab.py(reflow 两轮)、make_figures.py(配图)。实验台是一个二维八模高斯混合,噪声是 $N(0,I)$。选它的原因是:高斯混合经过任何高斯路径之后仍是高斯混合,于是 $m$、$e$、$v^*$ 全都有闭式解,不用训练也不用采样近似。 4.1 精确场:四条闭式 分量 $k$ 的中间方差是 $V_k(t)=\alpha_t^2 s_k^2+\sigma_t^2$($s_k$ 是该分量的标准差),后验责任 $r_k$ 由贝叶斯公式给出。三个量都是"责任加权的条件均值": def m(self, z, t): # E[x | z_t = z] a = self.path.alpha(t) r, V, diff = self._post(z, t) ex = self.data.mu[None,:,:] + (a * self.data.s**2 / V)[None,:,None] * diff return (r[:,:,None] * ex).sum(1) def e(self, z, t): # E[eps | z_t = z] sg = self.path.sigma(t) r, V, diff = self._post(z, t) ee = (sg / V)[None,:,None] * diff return (r[:,:,None] * ee).sum(1) def v(self, z, t): # 边缘速度场 return self.path.dalpha(t) * self.m(z, t) + self.path.dsigma(t) * self.e(z, t) 先验两个恒等式。第一个是 3.3 节的 $\alpha m+\sigma e=z$,第二个是把 $v^*$ 用另一种方式算一遍——直线路径下 3.3 节的表给出 $v^*=(z-m)/t$: === 恒等式自检:alpha*m + sigma*e 应等于 z(残差 ~1e-15)=== path t |a*m+s*e-z| |v-(da*m+ds*e)| linear 0.05 2.66e-15 0.00e+00 linear 0.95 1.78e-15 0.00e+00 cosine_vp 0.05 1.78e-15 0.00e+00 cosine_vp 0.95 1.78e-15 0.00e+00 === 直线路径下 v* 的两种算法是否一致:(z-m)/t 与 alpha' m + sigma' e === t=0.1 最大绝对差 = 1.332e-14 t=0.9 最大绝对差 = 1.776e-15 符号的物理含义都压到了 $10^{-15}$ 量级,说明推导和实现是同一件事。 4.2 换算表与放大倍数 把 3.3 节的换算写成代码并逐点验真值(注意残差随 $t\to 1$ 变大——那是除以 $(1-t)$ 的数值放大,不是推导错): def v_to_me(z, v, path, t): a, sg = path.alpha(t), path.sigma(t) da, ds = path.dalpha(t), path.dsigma(t) det = da * sg - ds * a # 直线路径恒等于 -1 m = (sg * v - ds * z) / det e = (da * z - a * v) / det return m, e path t v->m v->e eps->m eps->v linear 0.10 1.78e-15 1.33e-15 1.78e-15 1.78e-15 linear 0.90 1.22e-15 8.88e-16 1.30e-14 1.29e-14 linear 0.99 1.78e-15 1.78e-15 2.04e-13 2.03e-13 放大倍数不靠公式背,用数值扰动直接量:给 $m$ 加一个长度固定为 $0.01$ 的随机扰动,看三种预测误差范数各被放大多少;平方损失对应这些系数的平方。实测与解析式在小数点后四位完全一致($t=0.1$ 时噪声目标 $9.0000$ 对 $9.0000$、速度目标 $10.0000$ 对 $10.0000$),整条曲线就是图 1。 4.3 路径对比:NFE-误差曲线 误差尺子值得单独说一句。常用做法是"拿一条很高步数的参考解当真值",但那条参考解自己也有截断误差。这里利用 $p_t$ 仍是高斯混合这一点,取一组 RBF 特征 $\phi_c(z)=\exp(-\|z-c\|^2/(2h^2))$,它在高斯下的期望有闭式,于是参照侧没有任何采样噪声。生成样本一侧仍有有限样本误差:一批 4000 个真实样本对精确 $p_{\text{data}}$ 的偏差是 $0.00211$,这里只把它画成参考线,既不是误差下限,也不是显著性阈值。更换随机种子后,这个参考偏差也会变化。要判断小差异是否稳定,应对方法使用共同起点并重复多个种子。有限组 RBF 特征只是分布诊断,特征偏差小不能证明两个分布相同。 NFE 统计实际速度场评估次数:Euler 每步 1 次,RK2 每步 2 次。所以 NFE=8 时分别运行 8 步 Euler 或 4 步 RK2,表中按这一预算公平比较。 === NFE-误差:直线路径 vs VP 余弦路径(同一批起点、同一批随机数)=== NFE linear/euler linear/rk2 vp/euler vp/rk2 2 0.04074 0.04113 0.01047 0.04603 4 0.01491 0.00736 0.00562 0.00318 8 0.00802 0.00202 0.00329 0.00236 16 0.00466 0.00183 0.00234 0.00194 32 0.00307 0.00184 0.00199 0.00186 256 0.00194 0.00183 0.00184 0.00183 这张图要看什么:高 NFE 时,四条线都接近同一量级的小偏差,与有限样本误差并存;这不能单独证明场正确。低 NFE 的差异依赖路径和积分器:在 Euler、NFE=8 这一点,线性路径误差是 VP 的约 2.4 倍($0.00802$ 对 $0.00329$),不是“需要 2.4 倍步数”。要比较所需步数,应先规定相同误差阈值。RK2 在 NFE=2 时只有一个积分步,误差甚至大于 Euler;较高阶也不保证每个极低预算下都更好。 数值积分误差受速度场的空间与时间变化、时间网格和积分器共同影响,不能仅用“曲线弯”或错误的“两端速度为零”解释。时间 shift 也可能改变弯曲轨迹的积分精度;这里只能报告本次测试:Euler、NFE=8 时,shift=1、3、6 的误差分别为 $0.00802$、$0.00864$、$0.01843$,这两个非均匀网格没有带来改善。 4.4 Reflow:把弯的轨迹拉直 reflow 的实现出奇地短:用第 1 轮模型从噪声端积分到数据端,得到新配对,再训一轮。 def round1_pairs(n, seed=0): rng = base_rng(seed) data = make_data() z0 = data.sample(n, rng) # 数据端(t=0) z1 = rng.standard_normal((n, D)) # 噪声端(t=1),独立配对 return z0, z1 def generate_pairs(model, z1, n_steps=100): sched = make_schedule(n_steps) # RK2:100 步 = 200 次场评估 z0, _ = integrate(lambda z, t: model(z, t), z1, sched, "rk2") return z0 # 新配对就是 (z0, z1) 模型是 numpy 手搓的两层 128 宽 tanh MLP。第 1 轮在 10 万条独立配对上训练(每步从数据样本池抽样并重采独立噪声),训练 loss 收敛到 $2.1853$——按样本平方范数是 $4.3706$,而估计的不可约回归误差 $S_{pair}\cdot E\|z_1-z_0\|^2\approx4.3702$(未舍入计算;两个因子分别约为 $0.6229$、$7.0162$)。loss 接近这一估计,但有限采样下略低或略高都可能发生,不能证明模型已达精确最优;在分布内查询点上它与精确场的相对均方误差是 $1.28\%$。 然后是关键的一步——量两轮的归一化配对残差和 NFE: === 归一化配对回归残差(另用弧长/弦长量轨迹)=== 第 1 轮配对 + 精确场 S_pair = 0.6229 第 1 轮配对 + 第1轮模型 S_pair = 0.6289 第 2 轮配对 + 第2轮模型 S_pair = 0.0027 弧长/弦长:精确场 1.6499 | 第1轮 1.6463 | 第2轮 1.0027 === NFE-误差(直线路径;精确场 / 第1轮模型 / 第2轮模型)=== NFE 精确场 第1轮 第2轮 2 0.04012 0.04469 0.00402 4 0.01469 0.01789 0.00401 8 0.00785 0.00927 0.00402 16 0.00427 0.00563 0.00402 64 0.00182 0.00389 0.00403 这张图要看什么:左图是三条 NFE-误差曲线,绿色(reflow 后)从 NFE=2 起就是一条平线——步数不再是瓶颈;右图分别比较配对回归残差与实际轨迹弧长/弦长:配对残差第 1 轮用精确场、第 2 轮用重训模型;弧长柱均用对应训练模型。第 2 轮的弧长/弦长 $1.0027$,轨迹已经基本就是两点连线。另外注意第 2 轮配对的弦长平方从 $7.0162$ 掉到 $1.2410$:reflow 之后噪声端和数据端被强相关地配对了,这种配对改变也体现在实际轨迹变直上;不能把较小的配对残差本身称为曲率。 05. 工业级实现对照 以 diffusers 的 FlowMatchEulerDiscreteScheduler 为例(以 2026-09 的实现为准,scheduling_flow_match_euler_discrete.py),最小实现和生产实现的差别集中在四处。 一处一行更新。 step() 的默认分支就是速度空间里的显式 Euler: dt = sigma_next - sigma prev_sample = sample + dt * model_output model_output 被直接当作速度用,不做任何转换。和 04 节最小实现的 z = z - h * v 是同一行——diffusers 用 $\sigma$(就是本文的 $t$)当时间轴,$\sigma$ 从 1 递减到 0,所以步长是负的,方向自动正确。 换算真的写在生产代码里。 随机采样分支里有这一行: x0 = sample - current_sigma * model_output 这正是 3.3 节换算表的第一格 $m=z-t\,v$。也就是说"速度预测的模型可以零成本拿到干净图估计"这件事,在 diffusers 里是一行乘法。 模型评估点与最终样本落点不同。 默认噪声表的最后非零模型评估点约为 0.001,set_timesteps 仍在末尾追加 0。最后执行 sample + (0 - sigma) * model_output,样本实际到达 0,只是不再在 0 调用网络;这也等于直线路径的 $x_0$ 估计。二维 toy 的 linspace(1,0,...) 在循环中同样先评估当前非零时刻,再积分到目标,不是高维模型禁止的做法。 时间 shift 调整有效时间分配,生产模型常根据分辨率设置它。 静态 shift 的公式是 sigmas = self.shift * sigmas / (1 + (self.shift - 1) * sigmas) 而动态 shift 按序列长度插值:base_shift=0.5 对应 256 个 token,max_shift=1.15 对应 4096 个 token,中间用指数插值 $\exp(\mu)/(\exp(\mu)+(1/t-1))$ 过渡。动机是:更高分辨率往往增加空间相关信号的冗余,使同等加噪后信息仍较易恢复,实践中据此调整训练/采样的有效噪声分配;具体 shift 应遵循模型配置。04 节只验证了这个 toy 上的几个 shift 设置:它们没有降低对应预算下的误差,不能推广成 shift 对弯曲轨迹无效。 另外 SD3 论文在训练侧还做了一件事:时间 $t$ 不用均匀分布采,用 logit-normal(众数在中间)。对照图 1 就能理解为什么——不同时间采样相当于重新分配训练权重;SD3 直接回归 velocity,不能拿噪声参数化在高噪声端的弱信号直接解释其 logit-normal 选择。 06. 代价与边界 直线路径不是免费的午餐。 在这组实验的 Euler、NFE=8 处,线性路径的误差约为 VP 余弦路径的 2.4 倍;这个误差比不等于达到同一质量所需的步数比。需要诚实标注边界:这是一个二维、数据尺度(标准差 1.56)与噪声尺度(1)不匹配的 toy,不能据此断言"VP 路径普遍更省步数";但至少说明直线插值不保证少步积分准确。实际轨迹接近直线有助于理解 reflow 的效果,积分难度仍取决于沿轨迹的速度变化与数值格式。真实模型里路径选择还和训练分布、模型容量、蒸馏方案纠缠在一起,单独归因很难。 reflow 仍受教师分布与学习误差限制。 本次第 2 轮的高 NFE 偏差约为 $0.00403$;第 1 轮模型用于训练的 10 万个生成端点偏差为 $0.00328$。二者评估样本量不同,不能据此定量断言“误差几乎全部继承自教师”。在理论条件满足、回归和积分精确的理想 reflow 中,端点边缘分布保持为教师的分布;有限模型还会增加重训和积分误差。图 4 显示的是少步曲线被压平,不能当作重训自动改善真实数据拟合的证据。 reflow 的账单。 它需要用第 1 轮模型做一次全量生成,生成量要够训第 2 轮。其成本由配对生成量、教师步数和重训步数共同决定,不能固定说成翻倍,所以工程上要么只在最后的精调阶段做,要么直接换成少步蒸馏(知识树上的 step_distillation,待写)。 $\sigma\to0$ 的数值边界。 需要避免在端点进行会除以零的参数化转换;速度场 Euler 更新自身不含该除法。最后从 $\sigma>0$ 到 0 的 Euler 步正是 $x_0=x-\sigma v$,两种说法在直线路径下是同一操作。 07. 经典论文脉络 Lipman et al., 2022(arXiv:2210.02747)Flow Matching for Generative Modeling:用 simulation-free 条件回归训练速度场,避免训练期间反复数值求解 ODE;不是从 O(1) 复杂度降到 1,提出条件路径 + marginalization trick,是"用回归训速度场"的源头。 Liu et al., 2022(arXiv:2209.03003)Rectified Flow:从线性插值回归出发构造流,并提出 reflow。在论文的理想化条件下,reflow 具有直线度和凸运输成本的相关保证;本文 4.4 节的有限样本、有限模型实验用于观察趋势,不能替代理论条件。 Albergo & Vanden-Eijnden, 2022,Building Normalizing Flows with Stochastic Interpolants:通过随机端点的插值和回归目标构造确定性概率流。 Albergo, Boffi & Vanden-Eijnden, 2023,Stochastic Interpolants: A Unifying Framework for Flows and Diffusions:进一步引入可调噪声项,建立流与扩散过程的统一框架;应与上一项区分引用。 Karras et al., 2022(arXiv:2206.00364)EDM:不谈"流",把预条件(参数化)和时间调度当成独立的自由度来调。它把输出预条件、损失权重与采样调度分开设计;本文 3.3 节涉及其中的输出换算与相对权重,不能把换算恒等式当作训练行为完全等价。 Esser et al., 2024(arXiv:2403.03206)Stable Diffusion 3:把 rectified flow 推到大规模文生图,给出 logit-normal 时间采样与分辨率相关 shift——工业界从 DDPM 切换到流匹配的标志点。 08. 常见误解 误解一:"流匹配不是扩散模型,是另一套东西。" 同一族高斯路径,DDPM 是 $\alpha=\sqrt{\bar\alpha},\sigma=\sqrt{1-\bar\alpha}$ 的一支,Rectified Flow 是 $\alpha=1-t,\sigma=t$ 的一支。可以使用相同的网络架构;同一路径上的三种预测输出可以代数互转(端点需处理退化)(04 节残差 $10^{-15}$)。真正不同的是中间分布 $p_t$ 的形状和 loss 的加权,"两个流派"的说法遮住了这些可调的自由度。 误解二:"路径是直的,所以一两步就能出图。" 直的是训练时的插值线;采样走的是边缘场的积分曲线。实测 1-rectified 的弧长/弦长是 $1.65$,NFE=2 误差 $0.04469$。要两步出图,靠的是 reflow 之后的第 2 轮($0.00402$),不是第 1 轮的"直线路径"。 误解三:"reflow 之后质量和速度都变好了。" 只对了一半。步数-精度曲线确实被压平(NFE=2 就到位),但精度仍受教师样本分布与重训误差限制:第 2 轮在本次实验中停在约 $0.00403$,教师生成端点的参考偏差为 $0.00328$。理想 reflow 保持教师端点边缘,有限训练下则需另外评估分布偏差,不能保证质量自动提高。 误解四:"velocity 预测只是把回归目标换了个写法。" 真值层面是,loss 层面不是。$t=0.99$ 处噪声目标对干净图误差的放大倍数是 $0.0101$,速度目标是 $1.0101$,幅度系数差 100 倍、平方权重差 10000 倍;这只描述固定干净图误差的相对权重,不代表高噪声端一半训练样本都无用。 误解五:"流匹配需要知道边缘密度或 score。" 3.2 节的推导全程只用了条件期望和塔性质,训练目标里没有任何一项含 $p_t$。需要密度的是评估(比如算 NLL),不是训练。 09. 动手验证 三个可以自己跑的小实验,按顺序: 恒等式链:python fm_oracle.py。应该看到 $\alpha m+\sigma e=z$ 的残差在 $10^{-15}$ 量级,$v^*$ 的两种算法差 $10^{-14}$ 以内。如果你的实现里这两个数在 $10^{-6}$ 量级,大概率是后验责任没做 log-sum-exp。 放大倍数:python param_lab.py。看解析式与数值扰动两列是否完全一致,再看 $t=0.99$ 那一行——噪声目标 $0.0101$ 对速度目标 $1.0101$。 reflow 的残差与采样曲线:python reflow_lab.py(约两分钟,纯 numpy)。确认三件事:第 1 轮逐坐标平均 loss 乘以维数 2 后,接近估计的 $S_{pair}\cdot E\|z_1-z_0\|^2$;第 2 轮弧长/弦长 $\approx 1.00$;第 2 轮 NFE 曲线近乎平坦,但仍存在相对真实数据的偏差。 想改实验设置的话:make_data 里的 radius 控制数据尺度(默认 2.2)。改变它后重新运行 path_lab.py,比较不同路径在相同实际 NFE 下的误差。半径改变会同时影响数据协方差、模式间距和路径难度;不要把某一次误差比当作固定的步数收益。 10. 延伸阅读 前置:ddpm(VP 路径与噪声预测)、diffusion_math(前向反向推导与 score)、vae_elbo(另一条"回归式"生成目标)。 相邻:ddim_samplers(同一套 ODE 视角下的高阶格式与步长调度)、cfg(流匹配模型上的引导,公式形式几乎不变)。 后继:step_distillation(少步蒸馏,reflow 的工程替代品,待写);video_rl(已发布,Flow-GRPO 在流匹配模型上做在线强化学习,直接建立在这些速度场上)。 附录:完整代码 09 节用到的脚本全文如下(reflow_lab.py、path_lab.py、fm_oracle.py、param_lab.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 reflow_lab.py # -*- coding: utf-8 -*- """flow_matching 实验台(四):Reflow —— 用模型自己造的配对再训一轮。 Rectified Flow 最容易被转述错的一句话是「它的路径是直的」。 真的是直的,但直的是**训练时那条插值线** (1-t)x + t*eps; ODE 真正走出来的轨迹是边缘速度场 v*(z,t) = E[eps - x | z_t = z] 的积分曲线, 本例的积分轨迹会弯曲,弧长须把每一小段的欧式长度相加后再与弦长比较。 Reflow 做的事是:用第 1 轮训好的模型把噪声推到数据端,得到一批**配对** (z_1, z_0) = (起点噪声, 模型生成的终点),再拿这批配对重训一轮。 理想化 reflow 的性质有理论条件;本实验另外测量有限模型的轨迹长度比和采样误差。 这一份用 numpy 手搓一个小 MLP 当"模型",分别报告配对插值上的归一化回归残差 S_pair、实际 ODE 轨迹弧长/弦长和 NFE-误差。 """ import numpy as np from fm_oracle import ( D, base_rng, make_data, LinearPath, OracleFlow, make_rbf_centers, discrepancy, make_schedule, integrate, ) C, BW = make_rbf_centers() PATH = LinearPath() # ─────────────────── 一个小 MLP(numpy 手写) ─────────────────── class MLP: def __init__(self, sizes, rng): self.W, self.b = [], [] for i in range(len(sizes) - 1): lim = np.sqrt(6.0 / (sizes[i] + sizes[i + 1])) self.W.append(rng.uniform(-lim, lim, (sizes[i], sizes[i + 1]))) self.b.append(np.zeros(sizes[i + 1])) self.m = [np.zeros_like(w) for w in self.W] self.v = [np.zeros_like(w) for w in self.W] self.mb = [np.zeros_like(x) for x in self.b] self.vb = [np.zeros_like(x) for x in self.b] def forward(self, X): A = [X] H = X for i in range(len(self.W) - 1): H = np.tanh(H @ self.W[i] + self.b[i]) A.append(H) A.append(H @ self.W[-1] + self.b[-1]) return A[-1], A def step(self, X, Y, lr, tstep, beta1=0.9, beta2=0.999, eps=1e-8): pred, A = self.forward(X) g = 2.0 * (pred - Y) / pred.size for i in range(len(self.W) - 1, -1, -1): gW = A[i].T @ g gb = g.sum(0) self.m[i] = beta1 * self.m[i] + (1 - beta1) * gW self.v[i] = beta2 * self.v[i] + (1 - beta2) * (gW * gW) self.mb[i] = beta1 * self.mb[i] + (1 - beta1) * gb self.vb[i] = beta2 * self.vb[i] + (1 - beta2) * (gb * gb) mh = self.m[i] / (1 - beta1 ** tstep) vh = self.v[i] / (1 - beta2 ** tstep) # 先用本次前向的旧权重传播梯度,再更新参数。 g_prev = (g @ self.W[i].T) * (1 - A[i] ** 2) if i > 0 else None self.W[i] -= lr * mh / (np.sqrt(vh) + eps) bmh = self.mb[i] / (1 - beta1 ** tstep) bvh = self.vb[i] / (1 - beta2 ** tstep) self.b[i] -= lr * bmh / (np.sqrt(bvh) + eps) if i > 0: g = g_prev return float(((pred - Y) ** 2).mean()) def __call__(self, z, t): tt = np.full((z.shape[0], 1), float(t)) return self.forward(np.concatenate([z, tt], axis=1))[0] def train(pairs_z0, pairs_z1, n_step=4000, batch=512, lr=3e-3, seed=0, fresh=False): """在给定配对上训练速度场:目标 u = z_1 - z_0,z_t = (1-t) z_0 + t z_1。 fresh=True 时每一步现采一批新配对(第 1 轮的配对可以无限造), 避免模型把 2 万条配对背下来——背下来会让 loss 掉到"条件方差"以下, 看起来很美,其实场是有偏的。 """ rng = base_rng(seed) net = MLP([D + 1, 128, 128, D], np.random.default_rng(seed + 1)) n = pairs_z0.shape[0] losses = [] for s in range(1, n_step + 1): if fresh: idx = rng.integers(0, n, batch) z1 = rng.standard_normal((batch, D)) z0 = pairs_z0[idx] # 数据端样本池足够大,随机抽即视为新样本 else: idx = rng.integers(0, n, batch) z0 = pairs_z0[idx] z1 = pairs_z1[idx] t = rng.uniform(0.0, 1.0, (batch, 1)) zt = (1 - t) * z0 + t * z1 u = z1 - z0 X = np.concatenate([zt, t], axis=1) cur_lr = lr * (0.3 ** (s / n_step)) # 余弦退火换成简单指数退火 l = net.step(X, u, cur_lr, s) losses.append(l) return net, float(np.mean(losses[-200:])) # ─────────────────── 两轮的配对 ─────────────────── def round1_pairs(n, seed=0): """第 1 轮:数据与噪声独立配对(这就是 Rectified Flow 原始的训练配对)。""" rng = base_rng(seed) data = make_data() z0 = data.sample(n, rng) # 数据端(t=0) z1 = rng.standard_normal((n, D)) # 噪声端(t=1) return z0, z1 def generate_pairs(model, z1, n_steps=100): """用 RK2 的 n_steps 个积分步生成配对,实际场评估次数为 2*n_steps。""" sched = make_schedule(n_steps) z0, _ = integrate(lambda z, t: model(z, t), z1, sched, "rk2") return z0 # ─────────────────── 评价指标 ─────────────────── def straightness_of(pairs_z0, pairs_z1, vfun, nstep=32): """配对插值上的归一化回归残差;不是论文沿实际 ODE 轨迹定义的 straightness。""" ts = np.linspace(1.0, 0.0, nstep + 1)[:-1] dev = 0.0 for t in ts: zt = (1 - t) * pairs_z0 + t * pairs_z1 dev += float(np.mean(np.sum((vfun(zt, t) - (pairs_z1 - pairs_z0)) ** 2, axis=1))) dev /= len(ts) return dev / float(np.mean(np.sum((pairs_z1 - pairs_z0) ** 2, axis=1))) def arc_over_chord(pairs_z1, vfun, nstep=64): """真走一遍:从 z_1 出发积分到 t=0,量轨迹弧长与端点弦长之比。""" sched = make_schedule(nstep) z0, traj = integrate(vfun, pairs_z1, sched, "rk2") seg = np.linalg.norm(np.diff(traj, axis=0), axis=2).sum(0) chord = np.linalg.norm(z0 - pairs_z1, axis=1) return float(np.mean(seg / np.maximum(chord, 1e-12))), z0 def nfe_curve(pairs_z1, vfun, nfe_list, data): flow = OracleFlow(data, PATH) out = [] for nfe in nfe_list: sched = make_schedule(nfe) z0, _ = integrate(vfun, pairs_z1, sched, "euler") out.append(discrepancy(z0, 0.0, flow, C, BW)) return np.array(out) def field_error(model, flow, n=4000, seed=3): """训练出来的场与精确场差多少(相对均方)。 查询点取自真实边缘 p_t(与第 1 轮训练分布一致),并在五个时刻汇总平方误差。 """ rng = base_rng(seed) data = make_data() num = 0.0 den = 0.0 for t in (0.1, 0.3, 0.5, 0.7, 0.9): x = data.sample(n, rng) eps = rng.standard_normal((n, D)) zt = (1 - t) * x + t * eps v = flow.v(zt, t) num += float(np.mean(np.sum((model(zt, t) - v) ** 2, axis=1))) den += float(np.mean(np.sum(v ** 2, axis=1))) return num / den def chord_scale(pairs_z0, pairs_z1): """E||z_1 - z_0||^2,配对残差 S_pair 的分母。""" return float(np.mean(np.sum((pairs_z1 - pairs_z0) ** 2, axis=1))) def state(): data = make_data() flow = OracleFlow(data, PATH) z0_1, z1 = round1_pairs(100000, seed=0) net1, loss1 = train(z0_1, z1, n_step=6000, seed=0, fresh=True) z0_2 = generate_pairs(net1, z1, n_steps=100) net2, loss2 = train(z0_2, z1, n_step=6000, seed=0) nfe_list = [2, 4, 8, 16, 32, 64] eval_z1 = base_rng(999).standard_normal((4000, D)) out = { "nfe": np.array(nfe_list), "loss1": loss1, "loss2": loss2, "ferr1": field_error(net1, flow), "ferr2": field_error(net2, flow), "S1_oracle": straightness_of(z0_1, z1, lambda z, t: flow.v(z, t)), "S1_learned": straightness_of(z0_1, z1, lambda z, t: net1(z, t)), "S2_learned": straightness_of(z0_2, z1, lambda z, t: net2(z, t)), "arc1_oracle": arc_over_chord(z1, lambda z, t: flow.v(z, t))[0], "arc1_learned": arc_over_chord(z1, lambda z, t: net1(z, t))[0], "arc2_learned": arc_over_chord(z1, lambda z, t: net2(z, t))[0], "nfe_oracle": nfe_curve(eval_z1, lambda z, t: flow.v(z, t), nfe_list, data), "nfe_learned1": nfe_curve(eval_z1, lambda z, t: net1(z, t), nfe_list, data), "nfe_learned2": nfe_curve(eval_z1, lambda z, t: net2(z, t), nfe_list, data), } return out if __name__ == "__main__": data = make_data() flow = OracleFlow(data, PATH) print("=== 第 1 轮:独立配对上训练 ===") z0_1, z1 = round1_pairs(100000, seed=0) net1, loss1 = train(z0_1, z1, n_step=6000, seed=0, fresh=True) s1 = chord_scale(z0_1, z1) print(f" 训练 loss(末 200 步均值)= {loss1:.4f}(按样本平方范数是 {2*loss1:.4f})") s_pair_oracle = straightness_of(z0_1, z1, lambda z, t: flow.v(z, t)) print(f" 估计的不可约回归误差 = S_pair * E||z1-z0||^2 = " f"{s_pair_oracle:.4f} * {s1:.4f} = {s_pair_oracle*s1:.4f}") print(f" 学出来的场 vs 精确场,相对均方误差(on-distribution)= {field_error(net1, flow):.4f}") print() print("=== 用第 1 轮模型生成新配对 ===") z0_2 = generate_pairs(net1, z1, n_steps=100) print(f" 生成终点与精确 p_data 的特征偏差 = {discrepancy(z0_2, 0.0, flow, C, BW):.5f}" f"(另有一组 4000 个真实样本的参考偏差约 0.00211,非误差下限)") net2, loss2 = train(z0_2, z1, n_step=6000, seed=0) s2 = chord_scale(z0_2, z1) print(f" 第 2 轮训练 loss = {loss2:.4f}(按样本平方范数是 {2*loss2:.4f})") print(f" 第 2 轮弦长平方 E||z1-z0||^2 = {s2:.4f}(第 1 轮是 {s1:.4f},配对更紧了)") print() print("=== 归一化配对残差(另用弧长/弦长量轨迹)===") print(f" 第 1 轮配对 + 精确场 S_pair = {straightness_of(z0_1, z1, lambda z, t: flow.v(z, t)):.4f}") print(f" 第 1 轮配对 + 第1轮模型 S_pair = {straightness_of(z0_1, z1, lambda z, t: net1(z, t)):.4f}") print(f" 第 2 轮配对 + 第2轮模型 S_pair = {straightness_of(z0_2, z1, lambda z, t: net2(z, t)):.4f}") print(f" 弧长/弦长:精确场 {arc_over_chord(z1, lambda z, t: flow.v(z, t))[0]:.4f} | " f"第1轮 {arc_over_chord(z1, lambda z, t: net1(z, t))[0]:.4f} | " f"第2轮 {arc_over_chord(z1, lambda z, t: net2(z, t))[0]:.4f}") print() nfe_list = [2, 4, 8, 16, 32, 64] eval_z1 = base_rng(999).standard_normal((4000, D)) o = nfe_curve(eval_z1, lambda z, t: flow.v(z, t), nfe_list, data) a = nfe_curve(eval_z1, lambda z, t: net1(z, t), nfe_list, data) b = nfe_curve(eval_z1, lambda z, t: net2(z, t), nfe_list, data) print("=== NFE-误差(直线路径;精确场 / 第1轮模型 / 第2轮模型)===") print(" NFE 精确场 第1轮 第2轮") for i, n in enumerate(nfe_list): print(f"{n:>5}{o[i]:>10.5f}{a[i]:>10.5f}{b[i]:>10.5f}") path_lab.py # -*- coding: utf-8 -*- """flow_matching 实验台(三):路径的选择到底值多少 NFE。 直线路径(Rectified Flow)和 VP 余弦路径(DDPM 那一路)通向**同一个**目标分布, 中间分布不同,ODE 的轨迹与时间参数化也会不同。这一份把这件事量成三条曲线: 1. NFE-误差曲线:同一批起点,比较路径和积分器;Euler/RK2/RK4 每步分别评估 1/2/4 次场。 2. 欧式弧长 / 弦长,以及沿实际 ODE 轨迹的归一化速度偏差 S_traj; 后者对 Rectified Flow 的轨迹 straightness 再除以 E||Z_1-Z_0||^2, 不等于 reflow_lab 在原始配对插值上计算的 S_pair。 3. 时间 shift(SD3/FLUX 的做法)在直线路径上到底省不省步数。 误差尺子是 fm_oracle.discrepancy:拿精确边缘 p_t 的 RBF 特征矩当参照, 参照侧没有采样噪声,生成粒子一侧仍有有限样本误差。共享起点可减轻比较方差, 小差异是否稳定仍应更换随机种子验证;单次真实样本偏差不是硬下限。 """ import numpy as np from fm_oracle import ( D, base_rng, make_data, LinearPath, CosineVPPath, OracleFlow, make_rbf_centers, discrepancy, make_schedule, integrate, ) NPART = 4000 C, BW = make_rbf_centers() def start_particles(n, rng): """t=1 端的粒子:精确就是 N(0, I)。""" return rng.standard_normal((n, D)) def sampling_reference(data, n=NPART): """固定一次真实样本与精确 p_data 的偏差,仅作有限采样参考,不是误差下限。""" rng = base_rng(12345) x = data.sample(n, rng) flow = OracleFlow(data, LinearPath()) return discrepancy(x, 0.0, flow, C, BW) def run_sweep(path, method, nfe_list, npart=NPART): rng = base_rng() z1 = start_particles(npart, rng) flow = OracleFlow(make_data(), path) out = [] calls_per_step = {"euler": 1, "rk2": 2, "rk4": 4}[method] for nfe in nfe_list: if nfe < calls_per_step or nfe % calls_per_step: raise ValueError(f"{method}: NFE must be a positive multiple of {calls_per_step}") sched = make_schedule(nfe // calls_per_step) z0, _ = integrate(lambda z, t: flow.v(z, t), z1, sched, method) out.append(discrepancy(z0, 0.0, flow, C, BW)) return np.array(out) def straightness(path, npart=1500, nstep=400): """细步长积分出真实轨迹,量它离弦有多远。""" rng = base_rng(777) z1 = start_particles(npart, rng) flow = OracleFlow(make_data(), path) sched = make_schedule(nstep) z0, traj = integrate(lambda z, t: flow.v(z, t), z1, sched, "rk4") # traj: [nstep+1, npart, 2],index 0 是 t=1 arc = np.linalg.norm(np.diff(traj, axis=0), axis=2).sum(axis=0) # [npart],各段欧式长度相加 chord = np.linalg.norm(z0 - z1, axis=1) ratio = float(np.mean(arc / np.maximum(chord, 1e-12))) # S_traj = mean_t ||v(z_t,t) - (z_1 - z_0)||^2 / mean ||z_1 - z_0||^2 chord_dir = (z1 - z0)[None, :, :] # [1,npart,2] ts = sched[:-1] dev = 0.0 for i, t in enumerate(ts): v = flow.v(traj[i], t) dev += float(np.mean(np.sum((v - chord_dir[0]) ** 2, axis=1))) dev /= len(ts) scale = float(np.mean(np.sum((z1 - z0) ** 2, axis=1))) return ratio, dev / scale def shift_sweep(path, nfe_list, shifts, npart=NPART): rng = base_rng() z1 = start_particles(npart, rng) flow = OracleFlow(make_data(), path) res = {} for sh in shifts: errs = [] for nfe in nfe_list: sched = make_schedule(nfe, shift=sh) z0, _ = integrate(lambda z, t: flow.v(z, t), z1, sched, "euler") errs.append(discrepancy(z0, 0.0, flow, C, BW)) res[sh] = np.array(errs) return res def state(): """给画图脚本复用:返回与上面打印完全一致的数字。""" data = make_data() nfe_list = [2, 4, 8, 16, 32, 64, 128, 256] out = { "nfe": np.array(nfe_list), "sampling_reference": sampling_reference(data), "sweep": {}, "straight": {}, "shift": {}, } for path in (LinearPath(), CosineVPPath()): for method in ("euler", "rk2"): out["sweep"][(path.name, method)] = run_sweep(path, method, nfe_list) out["straight"][path.name] = straightness(path) out["shift"] = shift_sweep(LinearPath(), [4, 8, 16, 32, 64], [1.0, 3.0, 6.0]) out["shift_nfe"] = np.array([4, 8, 16, 32, 64]) return out if __name__ == "__main__": data = make_data() print(f"=== 单次采样参考:{NPART} 个真实样本 vs 精确 p_data 的偏差 = " f"{sampling_reference(data):.5f} ===") print("(这不是硬下限或显著性阈值;评估小差异需重复采样)") print() nfe_list = [2, 4, 8, 16, 32, 64, 128, 256] print("=== NFE-误差:直线路径 vs VP 余弦路径 ===") print(" NFE linear/euler linear/rk2 vp/euler vp/rk2") lin_e = run_sweep(LinearPath(), "euler", nfe_list) lin_r = run_sweep(LinearPath(), "rk2", nfe_list) vp_e = run_sweep(CosineVPPath(), "euler", nfe_list) vp_r = run_sweep(CosineVPPath(), "rk2", nfe_list) for i, n in enumerate(nfe_list): print(f"{n:>5}{lin_e[i]:>14.5f}{lin_r[i]:>13.5f}{vp_e[i]:>12.5f}{vp_r[i]:>11.5f}") print() print("=== 轨迹指标(400 步 RK4,即 1600 NFE;数值近似 ODE 轨迹)===") for path in (LinearPath(), CosineVPPath()): ratio, s = straightness(path) print(f" {path.label:<28} 弧长/弦长 = {ratio:.4f} 归一化 S_traj = {s:.4f}") print() print("=== 时间 shift(仅直线路径,Euler)===") res = shift_sweep(LinearPath(), [4, 8, 16, 32, 64], [1.0, 3.0, 6.0]) print(" NFE shift=1 shift=3 shift=6") for i, n in enumerate([4, 8, 16, 32, 64]): print(f"{n:>5}{res[1.0][i]:>11.5f}{res[3.0][i]:>11.5f}{res[6.0][i]:>11.5f}") fm_oracle.py # -*- coding: utf-8 -*- """flow_matching 实验台(一):任意高斯路径下的精确速度场。 这一份是整个文章所有数字的来源。思想是: 数据分布取二维高斯混合 p_data(K 个各向同性分量),噪声取 N(0, I)。 对任意「高斯路径」 z_t = alpha_t * x + sigma_t * eps(t=1 是纯噪声,t=0 是数据), p_t 仍然是高斯混合(分量均值 alpha*mu_k,方差 alpha^2 s_k^2 + sigma^2), 于是下面四样东西全都有闭式解,不需要训练、不需要采样近似: m(z,t) = E[x | z_t = z] 去噪均值(x0 预测的真值) e(z,t) = E[eps | z_t = z] 噪声均值(epsilon 预测的真值) v(z,t) = alpha' m + sigma' e 边缘速度场(速度预测的真值) score = grad log p_t(z) 有了这些,三种参数化(x0 / eps / v)可以逐点互相换算并验到浮点误差, ODE 的 NFE-误差曲线可以用「精确边缘 p_t」当尺子,不需要跑一条高精参考解。 误差尺子:取一组 RBF 特征 phi_c(z) = exp(-||z-c||^2 / (2h^2)), 它在高斯分布下的期望有闭式,于是 disc(粒子云, t) = mean_c | (1/N) sum_i phi_c(z_i) - E_{p_t}[phi_c] | 经验侧仍有有限样本误差,解析侧不额外引入采样噪声。单次真实样本参考值既不是硬地板, 也不是统计显著性阈值;有限组 RBF 特征矩一致也不能证明两个分布完全一致。 """ import numpy as np D = 2 # 二维,画图方便;公式与代码对任意维都成立 # ─────────────────────────── 随机流 ─────────────────────────── def base_rng(seed=20260928): """所有脚本与画图共用同一条随机流,保证正文数字与图上的数字一致。""" return np.random.default_rng(seed) # ─────────────────────────── 数据分布 ─────────────────────────── class GMM2D: """二维各向同性高斯混合。""" def __init__(self, means, stds, weights=None): self.mu = np.asarray(means, dtype=float) # [K, 2] self.s = np.asarray(stds, dtype=float) # [K] K = self.mu.shape[0] if weights is None: self.w = np.full(K, 1.0 / K) else: self.w = np.asarray(weights, dtype=float) self.w = self.w / self.w.sum() @property def K(self): return self.mu.shape[0] def sample(self, n, rng): k = rng.choice(self.K, size=n, p=self.w) return self.mu[k] + self.s[k][:, None] * rng.standard_normal((n, D)) def mean_cov(self): m = (self.w[:, None] * self.mu).sum(0) c = (self.w[:, None, None] * ( (self.s ** 2)[:, None, None] * np.eye(D)[None] + (self.mu - m)[:, :, None] * (self.mu - m)[:, None, :] )).sum(0) return m, c def make_data(K=8, radius=2.2, std=0.28): """8 个分量摆在半径 2.2 的圆上,每个分量标准差 0.28。""" ang = np.arange(K) * 2 * np.pi / K mu = np.stack([radius * np.cos(ang), radius * np.sin(ang)], axis=1) s = np.full(K, std) # 权重做成确定性的非均匀(1 + 0.35 cos),让混合更不像"一圈一样的点" w = 1.0 + 0.35 * np.cos(ang + 0.7) return GMM2D(mu, s, w / w.sum()) # ─────────────────────────── 路径 ─────────────────────────── class LinearPath: """Rectified Flow 的直线插值路径:z_t = (1-t) x + t eps。""" name = "linear" label = "直线路径(Rectified Flow)" def alpha(self, t): return 1.0 - t def sigma(self, t): return t def dalpha(self, t): return -1.0 def dsigma(self, t): return 1.0 class CosineVPPath: """方差保持(VP)路径:z_t = cos(pi t/2) x + sin(pi t/2) eps。 alpha^2 + sigma^2 = 1,也就是 DDPM 那一路 cosine schedule 的连续化版本。 """ name = "cosine_vp" label = "VP 余弦路径(DDPM 那一路)" def alpha(self, t): return np.cos(0.5 * np.pi * t) def sigma(self, t): return np.sin(0.5 * np.pi * t) def dalpha(self, t): return -0.5 * np.pi * np.sin(0.5 * np.pi * t) def dsigma(self, t): return 0.5 * np.pi * np.cos(0.5 * np.pi * t) # ─────────────────────────── 精确场 ─────────────────────────── class OracleFlow: """给定数据与路径后的精确边缘速度场。""" def __init__(self, data: GMM2D, path): self.data = data self.path = path # 后验responsibility r_k(z,t) 与分量方差 V_k(t) def _post(self, z, t): a = self.path.alpha(t) sg = self.path.sigma(t) V = a * a * self.data.s ** 2 + sg * sg # [K] diff = z[:, None, :] - a * self.data.mu[None, :, :] # [N,K,2] d2 = (diff ** 2).sum(-1) # [N,K] logp = -0.5 * d2 / V[None, :] - np.log(V)[None, :] + np.log(self.data.w)[None, :] logp = logp - logp.max(1, keepdims=True) r = np.exp(logp) r = r / r.sum(1, keepdims=True) return r, V, diff def m(self, z, t): """E[x | z_t = z],也就是 x0 预测的真值。""" a = self.path.alpha(t) r, V, diff = self._post(z, t) ex = self.data.mu[None, :, :] + (a * self.data.s ** 2 / V)[None, :, None] * diff return (r[:, :, None] * ex).sum(1) def e(self, z, t): """E[eps | z_t = z],也就是 epsilon 预测的真值。""" sg = self.path.sigma(t) r, V, diff = self._post(z, t) ee = (sg / V)[None, :, None] * diff return (r[:, :, None] * ee).sum(1) def v(self, z, t): """边缘速度场 v*(z,t) = alpha' m + sigma' e。""" da = self.path.dalpha(t) ds = self.path.dsigma(t) return da * self.m(z, t) + ds * self.e(z, t) def score(self, z, t): """grad log p_t(z)。""" r, V, diff = self._post(z, t) return -(r[:, :, None] * diff / V[None, :, None]).sum(1) def marginal(self, t): """p_t 的分量参数(仍是 GMM):均值 [K,2]、标准差 [K]、权重 [K]。""" a = self.path.alpha(t) sg = self.path.sigma(t) return a * self.data.mu, np.sqrt(a * a * self.data.s ** 2 + sg * sg), self.data.w # ─────────────────────────── 误差尺子 ─────────────────────────── def make_rbf_centers(n=9, extent=3.6, bw=0.9): g = np.linspace(-extent, extent, n) C = np.stack(np.meshgrid(g, g, indexing="ij"), axis=-1).reshape(-1, D) return C, bw def rbf_expectation(means, sds, weights, C, bw): """E_{p_t}[phi_c],phi_c(z)=exp(-||z-c||^2/(2 bw^2)),p_t 为各向同性 GMM。 单个高斯分量 N(m, v I) 下:E[phi_c] = (bw^2/(bw^2+v))^{d/2} * exp(-||m-c||^2/(2(bw^2+v))) """ v = sds ** 2 # [K] coef = (bw ** 2 / (bw ** 2 + v)) ** (D / 2.0) # [K] d2 = ((means[:, None, :] - C[None, :, :]) ** 2).sum(-1) # [K, M] val = coef[:, None] * np.exp(-0.5 * d2 / (bw ** 2 + v)[:, None]) return (weights[:, None] * val).sum(0) # [M] def discrepancy(z, t, flow: OracleFlow, C, bw): """粒子云 z 与精确边缘 p_t 的 RBF 特征矩偏差(越小越好)。""" mu_t, sd_t, w_t = flow.marginal(t) exact = rbf_expectation(mu_t, sd_t, w_t, C, bw) d2 = ((z[:, None, :] - C[None, :, :]) ** 2).sum(-1) emp = np.exp(-0.5 * d2 / bw ** 2).mean(0) return float(np.abs(emp - exact).mean()) # ─────────────────────────── 积分器 ─────────────────────────── def make_schedule(n_steps, shift=1.0): """含 n_steps 个积分区间(不是统一 NFE 预算)的时刻表。shift>1 是 SD3/FLUX 那套把时间往高噪声端推的做法。""" t = np.linspace(1.0, 0.0, n_steps + 1) if shift != 1.0: t = shift * t / (1.0 + (shift - 1.0) * t) return t def integrate(vfun, z, schedule, method="euler"): """从 t=1 走到 t=0。vfun(z, t) 返回对递增时间定义的 dz/dt;负步长负责反向积分。""" z = z.copy() traj = [z.copy()] for i in range(len(schedule) - 1): t = schedule[i] h = schedule[i] - schedule[i + 1] # >0 if method == "euler": dz = vfun(z, t) elif method == "rk2": zm = z - 0.5 * h * vfun(z, t) dz = vfun(zm, t - 0.5 * h) elif method == "rk4": k1 = vfun(z, t) k2 = vfun(z - 0.5 * h * k1, t - 0.5 * h) k3 = vfun(z - 0.5 * h * k2, t - 0.5 * h) k4 = vfun(z - h * k3, t - h) dz = (k1 + 2 * k2 + 2 * k3 + k4) / 6.0 else: raise ValueError(method) z = z - h * dz traj.append(z.copy()) return z, np.stack(traj) # traj: [S+1, N, 2] # ─────────────────────────── 自测 ─────────────────────────── def selfcheck(): rng = base_rng() data = make_data() out = [] for path in (LinearPath(), CosineVPPath()): flow = OracleFlow(data, path) z = rng.standard_normal((4000, D)) * 1.6 for t in (0.05, 0.25, 0.5, 0.75, 0.95): a = path.alpha(t) sg = path.sigma(t) m = flow.m(z, t) e = flow.e(z, t) v = flow.v(z, t) # 恒等式 1:alpha*m + sigma*e == z r1 = float(np.abs(a * m + sg * e - z).max()) # 恒等式 2:v == alpha' m + sigma' e(定义,顺手确认没写反导数) r2 = float(np.abs(v - (path.dalpha(t) * m + path.dsigma(t) * e)).max()) out.append((path.name, t, r1, r2)) return out if __name__ == "__main__": print("=== 恒等式自检:alpha*m + sigma*e 应等于 z(残差 ~1e-15)===") hdr = "path t |a*m+s*e-z| |v-(da*m+ds*e)|" print(hdr) for name, t, r1, r2 in selfcheck(): print(f"{name:<12}{t:>6.2f}{r1:>16.2e}{r2:>18.2e}") rng = base_rng() data = make_data() flow = OracleFlow(data, LinearPath()) print() print("=== 直线路径下 v* 的两种算法是否一致:(z-m)/t 与 alpha' m + sigma' e ===") for t in (0.1, 0.3, 0.5, 0.7, 0.9): z = rng.standard_normal((2000, D)) * 1.6 lhs = (z - flow.m(z, t)) / t rhs = flow.v(z, t) print(f" t={t:.1f} 最大绝对差 = {np.abs(lhs - rhs).max():.3e}") param_lab.py # -*- coding: utf-8 -*- """flow_matching 实验台(二):三种参数化(x0 / eps / v)的换算与 loss 权重。 要回答两个问题: 1. 固定同一条高斯路径和非退化时间,eps / x0 / v 输出能否代数换算? 可以换算同一路径上的输出,但不能据此把 VP 模型直接变成另一条 RF 路径的模型。 三种参数化的真值来自同一个去噪均值 m: z = alpha*m + sigma*e (路径定义,恒等) v = alpha'*m + sigma'*e (速度定义) 两式联立解出 m、e,就得到「速度 -> 噪声/干净图」的换算; 反过来也成立。下面把这些换算逐点验到 1e-14。 2. 输出可换算,为什么不同训练目标仍会产生差异? 因为「等价」指的是**真值**等价,**loss 不等价**:同一个 m 上的误差 delta, 在三种 loss 里被放大的倍数不同: x0 : 1 eps : alpha / sigma v : |alpha' - sigma' * alpha / sigma| 这里列的是误差幅度系数,平方损失权重还需平方。直线路径 t→1 时, eps 系数趋于 0,v 系数趋于 1;这比较固定 m 误差的相对权重, 不等于真实网络的所有梯度或质量。t=1 的 eps->m 换算本身退化。 """ import numpy as np from fm_oracle import ( D, base_rng, make_data, LinearPath, CosineVPPath, OracleFlow, ) # ─────────────────── 换算表的通用形式 ─────────────────── def v_to_me(z, v, path, t): """由 (z, v) 反解 (m, e)。 解 [[alpha', sigma'], [alpha, sigma]] @ [m, e] = [v, z] 行列式 det = alpha'*sigma - sigma'*alpha """ a, sg = path.alpha(t), path.sigma(t) da, ds = path.dalpha(t), path.dsigma(t) det = da * sg - ds * a m = (sg * v - ds * z) / det e = (da * z - a * v) / det return m, e def eps_to_me(z, eps_hat, path, t): """由 eps 预测反解 (m, e)。""" a, sg = path.alpha(t), path.sigma(t) m = (z - sg * eps_hat) / a return m, np.broadcast_to(eps_hat, m.shape).copy() def me_to_v(m, e, path, t): return path.dalpha(t) * m + path.dsigma(t) * e def check_conversions(): """真值层面三种参数化互转,误差应到浮点量级。""" rng = base_rng() data = make_data() rows = [] for path in (LinearPath(), CosineVPPath()): flow = OracleFlow(data, path) z = rng.standard_normal((3000, D)) * 1.6 for t in (0.1, 0.3, 0.5, 0.7, 0.9, 0.99): m = flow.m(z, t) e = flow.e(z, t) v = flow.v(z, t) # 真值 v -> (m,e) m2, e2 = v_to_me(z, v, path, t) # 真值 eps -> m -> v m3, _ = eps_to_me(z, e, path, t) v3 = me_to_v(m3, e, path, t) rows.append(( path.name, t, float(np.abs(m2 - m).max()), float(np.abs(e2 - e).max()), float(np.abs(m3 - m).max()), float(np.abs(v3 - v).max()), )) return rows # ─────────────────── 放大倍数 ─────────────────── def weights(path, ts): """返回三种参数化下「m 上的单位误差」被放大的倍数(振幅,非平方)。""" wx = np.ones_like(ts) we, wv = [], [] for t in ts: a, sg = path.alpha(t), path.sigma(t) da, ds = path.dalpha(t), path.dsigma(t) we.append(a / sg) wv.append(abs(da - ds * a / sg)) return wx, np.array(we), np.array(wv) def numeric_weight_check(): """不靠公式,直接用数值扰动验证放大倍数:给 m 加一个固定扰动,看 loss 变化。""" rng = base_rng() data = make_data() flow = OracleFlow(data, LinearPath()) path = LinearPath() z = rng.standard_normal((4000, D)) * 1.6 out = [] for t in (0.1, 0.3, 0.5, 0.7, 0.9, 0.99): m = flow.m(z, t) e = flow.e(z, t) v = flow.v(z, t) delta = rng.standard_normal((4000, D)) delta = delta / np.linalg.norm(delta) * 0.01 # 固定长度 0.01 的扰动 mp = m + delta ep = (z - path.alpha(t) * mp) / path.sigma(t) vp = me_to_v(mp, ep, path, t) r_eps = np.linalg.norm(ep - e) / np.linalg.norm(delta) r_v = np.linalg.norm(vp - v) / np.linalg.norm(delta) _, we, wv = weights(path, np.array([t])) out.append((t, r_eps, float(we[0]), r_v, float(wv[0]))) return out if __name__ == "__main__": print("=== 换算残差(真值层面互转,应为 1e-14 量级)===") print("path t v->m v->e eps->m eps->v") for name, t, r1, r2, r3, r4 in check_conversions(): print(f"{name:<11}{t:>5.2f}{r1:>12.2e}{r2:>12.2e}{r3:>12.2e}{r4:>12.2e}") print() print("=== 放大倍数:解析式 vs 数值扰动(直线路径)===") print(" t eps实测 eps解析 v实测 v解析") for t, re_, we_, rv_, wv_ in numeric_weight_check(): print(f"{t:>5.2f}{re_:>10.4f}{we_:>10.4f}{rv_:>10.4f}{wv_:>10.4f}") print() print("=== 放大倍数随 t 的变化(直线路径 vs VP 余弦路径)===") ts = np.array([0.01, 0.05, 0.1, 0.2, 0.4, 0.6, 0.8, 0.9, 0.95, 0.99]) for path in (LinearPath(), CosineVPPath()): wx, we, wv = weights(path, ts) print(f"-- {path.label}") print(" t x0 eps v") for i, t in enumerate(ts): print(f"{t:>7.2f}{wx[i]:>8.3f}{we[i]:>10.3f}{wv[i]:>10.3f}") make_figures.py # -*- coding: utf-8 -*- """画配图。数字全部来自同目录的实验脚本,不另算一遍。 运行: python make_figures.py [--only 图名] """ import os import sys import numpy as np import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) from fm_oracle import ( # noqa: E402 D, base_rng, make_data, LinearPath, CosineVPPath, OracleFlow, make_schedule, integrate, ) from param_lab import weights # noqa: E402 import path_lab # noqa: E402 import reflow_lab # noqa: E402 HERE = os.path.dirname(os.path.abspath(__file__)) FIGDIR = os.path.join(os.path.dirname(HERE), "figures") os.makedirs(FIGDIR, exist_ok=True) C0 = "#2f4b7c" C1 = "#d45087" C2 = "#f0a35e" C3 = "#4c9f70" CGREY = "#8a8a8a" plt.rcParams["font.sans-serif"] = ["Heiti TC", "Arial Unicode MS", "DejaVu Sans"] plt.rcParams["axes.unicode_minus"] = False def _save(fig, name): p = os.path.join(FIGDIR, name) fig.savefig(p, dpi=150, bbox_inches="tight", facecolor="white") plt.close(fig) print(f" [ok] {name} ({os.path.getsize(p)} bytes)") # ─────────────── 图 1:三种参数化的放大倍数 ─────────────── def fig_weights(): ts = np.linspace(0.01, 0.999, 400) wx, we, wv = weights(LinearPath(), ts) fig, ax = plt.subplots(figsize=(7.4, 4.2)) ax.plot(ts, wv, color=C1, lw=2.4, label="velocity 参数化 v") ax.plot(ts, we, color=C0, lw=2.4, label="噪声参数化 eps") ax.plot(ts, wx, color=CGREY, lw=1.8, ls="--", label="干净图参数化 x0") ax.set_yscale("log") ax.set_xlabel("时间 t(0 = 数据端,1 = 纯噪声端)") ax.set_ylabel("同一份误差被放大的倍数(对数轴)") ax.set_title("图 1:预测误差幅度系数(平方损失权重为其平方)") ax.axvline(0.99, color="#cccccc", lw=1, ls=":") ax.annotate("t=0.99 处 eps 系数为 0.010,\nvelocity 还有 1.010", xy=(0.99, 1.0), xytext=(0.55, 0.35), fontsize=10, color="#333333", arrowprops=dict(arrowstyle="->", color="#999999", lw=1)) ax.legend(loc="upper right", fontsize=10) ax.grid(alpha=0.25) _save(fig, "fig1_param_weights.png") # ─────────────── 图 2:两条路径的真实 ODE 轨迹 ─────────────── def fig_trajectories(): data = make_data() rng = base_rng(4242) z1 = rng.standard_normal((600, D)) xs = data.sample(1200, base_rng(11)) fig, axes = plt.subplots(1, 2, figsize=(12.0, 5.4), sharex=True, sharey=True) sched = make_schedule(300) for ax, path in zip(axes, (LinearPath(), CosineVPPath())): flow = OracleFlow(data, path) z0, traj = integrate(lambda z, t: flow.v(z, t), z1, sched, "rk4") ax.scatter(xs[:, 0], xs[:, 1], s=8, color="#dddddd", label="数据样本(8 个模式)") for i in range(14): ax.plot(traj[:, i, 0], traj[:, i, 1], color=C1, lw=1.3, alpha=0.9) ax.scatter(z1[:14, 0], z1[:14, 1], s=34, color=C0, zorder=5, label="起点(纯噪声)") ax.scatter(z0[:14, 0], z0[:14, 1], s=34, color=C3, marker="s", zorder=5, label="终点(生成)") ax.set_title(path.label) ax.set_xlabel("dim 1") ax.grid(alpha=0.2) axes[0].set_ylabel("dim 2") axes[0].legend(fontsize=9, loc="upper left") fig.suptitle("图 2:同一批起点,两条路径走出来的 ODE 轨迹(300 步 RK4)", fontsize=13) fig.tight_layout() _save(fig, "fig2_trajectories.png") # ─────────────── 图 3:NFE-误差 ─────────────── def fig_nfe(): nfe_list = [2, 4, 8, 16, 32, 64, 128, 256] floor = path_lab.sampling_reference(make_data()) cur = { "linear/euler": path_lab.run_sweep(LinearPath(), "euler", nfe_list), "linear/rk2": path_lab.run_sweep(LinearPath(), "rk2", nfe_list), "vp/euler": path_lab.run_sweep(CosineVPPath(), "euler", nfe_list), "vp/rk2": path_lab.run_sweep(CosineVPPath(), "rk2", nfe_list), } fig, ax = plt.subplots(figsize=(7.6, 4.6)) for (k, v), col, mk in zip(cur.items(), [C1, C0, C2, C3], ["o", "o", "s", "s"]): ax.loglog(nfe_list, v, marker=mk, color=col, lw=2, label=k) ax.loglog(nfe_list, [floor] * len(nfe_list), color=CGREY, ls="--", lw=1.6, label="单次采样参考(4000 个真实样本)") ax.set_xlabel("NFE(模型评估次数)") ax.set_ylabel("终点分布与精确 p_data 的偏差") ax.set_title("图 3:1-rectified 的直线路径 vs VP 余弦路径") ax.grid(alpha=0.25, which="both") ax.legend(fontsize=10) _save(fig, "fig3_nfe_paths.png") return {"nfe": nfe_list, "curves": {k: v.tolist() for k, v in cur.items()}} # ─────────────── 图 4:Reflow ─────────────── def fig_reflow(): data = make_data() flow = OracleFlow(data, LinearPath()) z0_1, z1 = reflow_lab.round1_pairs(100000, seed=0) net1, _ = reflow_lab.train(z0_1, z1, n_step=6000, seed=0, fresh=True) z0_2 = reflow_lab.generate_pairs(net1, z1, n_steps=100) net2, _ = reflow_lab.train(z0_2, z1, n_step=6000, seed=0) nfe_list = [2, 4, 8, 16, 32, 64] eval_z1 = base_rng(999).standard_normal((4000, D)) o = reflow_lab.nfe_curve(eval_z1, lambda z, t: flow.v(z, t), nfe_list, data) a = reflow_lab.nfe_curve(eval_z1, lambda z, t: net1(z, t), nfe_list, data) b = reflow_lab.nfe_curve(eval_z1, lambda z, t: net2(z, t), nfe_list, data) floor = path_lab.sampling_reference(data) arc1 = reflow_lab.arc_over_chord(z1, lambda z, t: net1(z, t))[0] arc2 = reflow_lab.arc_over_chord(z1, lambda z, t: net2(z, t))[0] s1 = reflow_lab.straightness_of(z0_1, z1, lambda z, t: flow.v(z, t)) s2 = reflow_lab.straightness_of(z0_2, z1, lambda z, t: net2(z, t)) fig, axes = plt.subplots(1, 2, figsize=(12.0, 4.8)) ax = axes[0] ax.loglog(nfe_list, o, marker="o", color=CGREY, lw=2, label="精确场(数值积分对照)") ax.loglog(nfe_list, a, marker="o", color=C1, lw=2, label="第 1 轮模型") ax.loglog(nfe_list, b, marker="s", color=C3, lw=2, label="第 2 轮模型(reflow 后)") ax.loglog(nfe_list, [floor] * len(nfe_list), color="#bbbbbb", ls="--", lw=1.5, label="单次采样参考(非下限)") ax.set_xlabel("NFE") ax.set_ylabel("终点偏差") ax.set_title("本实验 reflow 后:NFE≥2 的偏差变化很小") ax.grid(alpha=0.25, which="both") ax.legend(fontsize=9) ax = axes[1] names = ["第 1 轮", "第 2 轮"] x = np.arange(2) ax.bar(x - 0.18, [arc1, arc2], width=0.34, color=C0, label="弧长/弦长(1 = 完全直线)") ax.bar(x + 0.18, [s1, s2], width=0.34, color=C2, label="归一化配对残差 S_pair") ax.set_xticks(x) ax.set_xticklabels(names) ax.set_yscale("log") ax.set_ylabel("配对残差与轨迹长度比(对数轴)") ax.set_title("两轮之间:轨迹从弯的变成直的") ax.grid(alpha=0.25, axis="y") ax.legend(fontsize=9) for i, (av, sv) in enumerate([(arc1, s1), (arc2, s2)]): ax.text(i - 0.18, av * 1.08, f"{av:.4f}", ha="center", fontsize=9) ax.text(i + 0.18, sv * 1.08, f"{sv:.4f}", ha="center", fontsize=9) fig.suptitle("图 4:Reflow 的两轮对比(同一个直线路径、同一批起点)", fontsize=13) fig.tight_layout() _save(fig, "fig4_reflow.png") return {"nfe1": a.tolist(), "nfe2": b.tolist(), "S1": s1, "S2": s2, "arc1": arc1, "arc2": arc2} if __name__ == "__main__": only = sys.argv[2] if len(sys.argv) > 2 and sys.argv[1] == "--only" else None jobs = { "weights": fig_weights, "trajectories": fig_trajectories, "nfe": fig_nfe, "reflow": fig_reflow, } for k, fn in jobs.items(): if only and k != only: continue print(f"[draw] {k}") fn() print("done")
2026年09月28日
3 阅读
0 评论
0 点赞
2026-09-28
AIGC 基本功|分类器无关引导 CFG 的代价与调法-CFG
分类器无关引导 CFG 的代价与调法 所属方向:生成范式 | 难度:进阶 | 前置知识:DDPM 训练目标与采样流程($\varepsilon$ 预测器、$\bar\alpha_t$ 与 DDIM 那一步怎么走都在那篇推过,这里直接用结论) 关键词:CFG、guidance scale、噪声预测外推、倾斜分布、guidance interval、guidance rescale、autoguidance 01. 为什么需要它 我第一次认真算 CFG 的账,是被一句「CFG 不增加参数,所以是免费的」气到的。它不增加参数,但它让每一步的网络前向从一次变成两次——这个 2× 是定义级的,不是工程上的估计。后面会把账算到 MAC 和显存字节上,先看它到底换来了什么。 我搭了一个二维高斯混合的玩具:4 个高斯分量,2 个类别,两类故意重叠(类心只差 1.3,标准差 0.6~1.0),每个类里一个紧分量(方差 0.36)一个松分量(方差 1.00)。选它的理由是:高斯混合的 MMSE 去噪器 $E[x_0|x_t]$ 有闭式解,条件分支和无条件分支都是精确的。也就是说,这个玩具里没有任何训练误差,可把训练误差排除,但仍要控制有限步采样、有限样本与数值积分误差。 然后我故意换上一个弱去噪器——把整簇拟合成一个高斯的线性维纳滤波。它有误差,而且误差随 $t$ 变,这才有资格代表真实 UNet。结果是这样的: $w$ 生成方差 / 真实条件方差 软纯度 1.0 0.5847 0.7379 2.0 0.4012 0.8787 3.0 0.2751 0.9312 5.0 0.1290 0.9553 7.5 0.0499 0.9648 15.0 0.0028 0.9696 $w$ 从 1 拉到 15,纯度只涨了 0.2317,多样性塌到真实值的 0.28%——画面上就是所有样本缩成一个点。这就是真实世界里「CFG 调过了就一片死板、颜色发焦」的原型。注意 $w=1$ 那一行的方差比是 0.5847 而不是 1:弱去噪器本身就在过平滑,CFG 不是来修它的,CFG 是在过平滑的基础上再往目标类上推。 把弱去噪器换回精确去噪器,故事变了,但没变好: $w$ 软纯度 硬纯度 多样性比 典型度 (nats) 类心偏移 真实条件样本(参照) 0.7388 0.7834 1.0000 +0.0000 0.0000 0.0 0.5001 0.5008 1.2512 +0.0081 0.6651 1.0 0.7409 0.7852 0.9876 +0.0030 0.0075 2.0 0.8770 0.9545 0.8596 −0.1995 0.4546 3.0 0.9245 0.9969 0.8193 −0.4974 0.7721 5.0 0.9503 1.0000 0.8042 −1.2480 1.2615 7.5 0.9636 1.0000 0.7933 −2.3042 1.7809 10.0 0.9744 1.0000 0.7603 −3.3963 2.2545 15.0 0.9905 1.0000 0.6716 −5.5530 3.0681 25.0 0.9985 1.0000 0.6033 −9.7613 4.2532 先看第一行和 $w=1$ 那一行:它们几乎完全一样(软纯度 0.7388 对 0.7409,典型度 +0.0000 对 +0.0030,类心偏移 0.0000 对 0.0075)。在去噪器精确的前提下,$w=1$ 采出来的就是真实条件分布 $p(x|c)$——这是对的,因为无条件分支和条件分支都是闭式解,没有误差可修。 那么 $w>1$ 在买什么?看 $w=7.5$:软纯度从 0.7409 涨到 0.9636,但典型度掉到 −2.3042 nats,类心偏移 1.78。典型度是「生成样本的平均 $\log p_{\text{data}}$」减去「真实样本的平均 $\log p_{\text{data}}$」,负值意味着生成样本平均落在较低的数据密度区;高斯混合的支撑是整个空间,不能称为“支撑之外”。也就是说:CFG 不是在把分布修得更准,它是在换一个目标分布,而且换过去的那一边偏离高密度数据区域。$w=25$ 时典型度 −9.7613、类心偏移 4.25,画面上就是过饱和、结构崩坏。 这一篇要讲的就是这个交易:用两倍算力,买一个明确的偏离。讲清楚代价的构成(算力、误差放大、幅度膨胀),才能讲清楚三个可调旋钮(强度 $w$、引导区间、rescale)分别在动哪一根杠杆。 02. 最小可用理解 三句话: 机制:每一步跑两遍网络,一次带条件 $c$、一次带空条件,然后把无条件预测沿「条件减无条件」的方向外推 $w$ 倍:$e_{\text{guided}}=e_{\text{un}}+w(e_{\text{c}}-e_{\text{un}})$。 成本:常规双分支实现的等价样本级网络计算约翻倍;本文卷积账本在 batch 维翻倍、权重共享的口径下,MAC 与逐层张量字节和都翻倍——$1.898\times10^{11}\to3.796\times10^{11}$ MAC、$0.13\ \text{GB}\to0.25\ \text{GB}$,比值都是 2.0000。只要 $w>1$ 就是这个价,跟 $w$ 是 2 还是 30 无关。 效果:每个噪声时刻的组合 score 可形式上写成倾斜密度的 score,但通常不能保证最终输出服从 $p(x|c)^{w}p(x)^{1-w}$。$w$ 越大越像目标类、越不像真实数据;$w=1$ 回到普通条件采样,$w=0$ 回到无条件采样。 这张图要看什么:左轴两条实线是软纯度(蓝,$E[p(c|x)]$)和多样性比(绿,生成方差 / 真实条件方差),右轴两条虚线是典型度(红)和类心偏移(紫)。四条线在 $w\approx2\sim3$ 附近同时拐弯:纯度在那之前涨得最快(0.7409→0.9245),之后收益迅速变平(0.9245→0.9985 用了 22 个 $w$);而典型度和类心偏移是没有平台期的,一路线性往下掉。灰色水平点线是真实条件样本自己的软纯度 0.7388——$w>1$ 的曲线整体在它上方,这就是「买纯度」的字面意思。竖线两条:灰色实线是 $w=1$($0<w<1$ 时在两分支之间插值,$w<0$ 才沿反方向外推),红色点划线是 $w=7.5$(SD 系默认值)。所以「$w=7.5$ 是常用默认值」这句话的实质是:把后面那段边际收益极低、但代价线性增长的区间也一并买了。 03. 数学推导 3.1 符号与定义 DDPM 的前向过程写成 $x_t=\sqrt{\bar\alpha_t}x_0+\sqrt{1-\bar\alpha_t}\,\varepsilon$,网络学的是 $\varepsilon_\theta(x_t,t,c)\approx\varepsilon$。记 $$e_{\text{c}}=\varepsilon_\theta(x_t,t,c),\qquad e_{\text{un}}=\varepsilon_\theta(x_t,t,\varnothing)$$ 其中 $\varnothing$ 是空条件(训练时按一定概率丢弃条件,让同一个权重同时充当两个分支)。CFG 的输出是 $$e_{\text{guided}}=e_{\text{un}}+w\,(e_{\text{c}}-e_{\text{un}})=w\,e_{\text{c}}+(1-w)\,e_{\text{un}}$$ 这个改写值得停一下:它是两支预测的仿射组合,系数和为 1,但其中一个系数为负。$w=1$ 时退化成纯条件分支,$w=0$ 时退化成纯无条件分支,$w>1$ 时无条件分支的系数 $1-w<0$,也就是把「条件相对无条件的优势」往远处推。 常用共享权重实现需要用条件丢弃训练空条件分支;定义上也可以分别训练条件和无条件模型,不是必须共享权重。做法是训练时按固定概率把条件替换成空条件(文生图里常用 10%~20%,具体比例随模型配置变化,不能把 10% 当成全系列统一默认),损失函数完全不变,只是输入换了一个占位 embedding。所以 CFG 不需要额外参数、不需要第二个网络,代价全部发生在推理侧——这一点是理解全文的钥匙:它把成本从「训练时一次」挪到了「推理时每次」。也因此,CFG 的强度是采样时才决定的超参,同一个 checkpoint 可以在 $w=3$ 和 $w=15$ 之间随便切。 3.2 它和「倾斜分布」是什么关系 用 score 的语言会更清楚。$\varepsilon$ 预测与 score 只差一个缩放:$\varepsilon_\theta=-\sqrt{1-\bar\alpha_t}\,\nabla_{x_t}\log p_\theta(x_t|c)$。把这个关系代进 3.1 的式子,$\sqrt{1-\bar\alpha_t}$ 整体提出来抵消,得到 $$\tilde{s}(x_t|c)=\nabla\log p(x_t|\varnothing)+w\big(\nabla\log p(x_t|c)-\nabla\log p(x_t|\varnothing)\big)=(1-w)\nabla\log p(x_t)+w\,\nabla\log p(x_t|c)$$ 而一个倾斜分布 $\tilde p(x|c)\propto p(x|c)^{w}p(x)^{1-w}$ 的 score 恰好就是 $$\nabla\log\tilde p(x|c)=w\,\nabla\log p(x|c)+(1-w)\,\nabla\log p(x)$$ 两式逐字相同。所以教科书里那句「CFG 采样自 $p(x|c)^w p(x)^{1-w}$」,在 score 层面是恒等的。 问题出在下一步:这个恒等式只保证每一步的 score 是对的,不保证采样出来的分布是 $\tilde p$。要让整条链落在 $\tilde p$ 上,需要每一步的边际 $p(x_t|c)$ 也被同样地倾斜,而这一点并不由上式推出——真实链上的 $x_t$ 来自上一步的输出,分布已经不是 $\tilde p$ 的边际了。$w$ 越大,这个偏差越大。第 06 节会用交叉熵把它量化:在 $w=7.5$ 时,CFG 采出的分布相对倾斜目标的交叉熵是 3.6926,而倾斜目标自己对自己的交叉熵是 2.2258,差了 1.47 nats——不是小偏差,是两个不同的分布。 几何上看得更直接: 这张图要看什么:三张等高线在同一个坐标系(网格 $[-4.2,4.2]^2$)下并排——左是真实条件分布 $p(x|c)$,中是倾斜目标 $p(x|c)^{7.5}p(x)^{-6.5}$,右是 CFG 实际采出的分布。三个面板里的灰色细线是同一组无条件数据密度等值线,用它当标尺。黑色加号是 4 个高斯分量的类心,蓝色圆点是真实条件均值,红色菱形是生成分布均值,箭头从前者指到后者(长度 1.78,而数据每维标准差约 0.91——偏了将近两个标准差)。要看的是中图和右图的差别:倾斜目标仍然贴着数据密度的等高线走,只是把权重在已有支撑上重新分配;而 CFG 采出的分布已经更多质量被推向低密度区域。这就是「偏离高密度数据区域」的几何版本,也是「过曝」这两个字的字面意思。 3.3 误差放大:为什么最坏上界是 $2w-1$ 设真实的条件/无条件预测为 $e^\star_{\text{c}}$ 与 $e^\star_{\text{un}}$,模型误差为 $\delta_{\text{c}}=e_{\text{c}}-e^\star_{\text{c}}$、$\delta_{\text{un}}=e_{\text{un}}-e^\star_{\text{un}}$。引导后预测相对「同样加权过的真值」$w e^\star_{\text{c}}+(1-w)e^\star_{\text{un}}$ 的误差是 $$\delta_{\text{guided}}=w\,\delta_{\text{c}}+(1-w)\,\delta_{\text{un}}$$ 取范数并放缩: $$\|\delta_{\text{guided}}\|\le w\|\delta_{\text{c}}\|+(w-1)\|\delta_{\text{un}}\|\le(2w-1)\max\big(\|\delta_{\text{c}}\|,\|\delta_{\text{un}}\|\big)$$ 注意这里用的是 $|1-w|=w-1$,不是 $1-w$。系数 $1-w$ 是负的这件事,正是代价的来源:三角不等式给出最坏上界;负系数既可能使误差叠加,也可能使相关误差抵消。$w=7.5$ 时上界是 14 倍,$w=15$ 时是 29 倍。实测远小于这个界(第 06 节给数字),因为两支误差高度相关,但本 toy 的这些探测点上呈单调放大,但不是对任意两支误差都成立的定理。 3.4 幅度膨胀:过曝的机制 即使两支都精确,引导后的预测幅度也会膨胀。看 $x_0$ 的重建:$x_0=(x_t-\sqrt{1-\bar\alpha_t}\,e_{\text{guided}})/\sqrt{\bar\alpha_t}$,它是 $e_{\text{guided}}$ 的仿射函数。$e_{\text{guided}}$ 的幅度一大,重建的 $x_0$ 就被推到远离数据中心的位置——这就是类心偏移的来源,也是画面上「过曝」的直接机制。实测两支预测的逐元素相关系数 $\rho=0.9928$(高度相关),所以膨胀不是来自两支的水平差异,而是来自差值方向 $e_{\text{c}}-e_{\text{un}}$ 被乘了 $w$ 之后叠加在一个本来就很大的共同分量上。第 06 节会给出 $\text{std}(e_{\text{guided}})/\text{std}(e_{\text{c}})$ 随 $w$ 的曲线。 04. 代码实现 核心是「精确去噪器 + DDIM」。高斯混合下 $E[x_0|x_t]$ 有闭式解:分量内部是高斯的,所以 $$E[x_0|x_t,k]=\mu_k+\frac{\bar\alpha_t^{1/2}\,\sigma_k^2}{\bar\alpha_t\sigma_k^2+(1-\bar\alpha_t)}\big(x_t-\bar\alpha_t^{1/2}\mu_k\big)$$ 再按分量后验 $r_k=p(k|x_t)$ 加权。代码就是这一行公式: def x0_hat(X, a_bar, logw): s = np.sqrt(a_bar) v = a_bar * VAR + (1.0 - a_bar) gain = VAR * s / v r = posterior(X, a_bar, logw) per = MU[None, :, :] + gain[None, :, None] * (X[:, None, :] - s * MU[None, :, :]) return (r[:, :, None] * per).sum(axis=1) 条件分支和无条件分支共用这个函数,只换先验权重:无条件用全混合的分量先验 $\pi_k$,条件用「只保留目标类、类内重新归一化」的先验。这是 CFG 训练方式的最小抽象——不是两个网络,是同一个网络喂两个条件。 这里有个实现上的坑值得单独说:类内重新归一化之后,类外分量的先验变成 0,$\log 0$ 会直接炸。所以所有先验都要先过一遍 safe_log,把非正的权重映到 $-\infty$ 而不是让 numpy 抛 RuntimeWarning。这个坑在写倾斜分布那段还会再咬一次——$\log p_{\text{tilde}}$ 的归一化常数要在网格上做数值积分,网格范围取窄了(比如只取 $\pm3.6$)会把分布的尾巴切掉,归一化常数偏小,后面所有交叉熵都跟着错。本文用的是 $\pm9$ 的 420 点网格。 $\varepsilon$ 预测和 DDIM 的一步: def eps_hat(X, t, logw): a = abar(t) return (X - np.sqrt(a) * x0_hat(X, a, logw)) / np.sqrt(1.0 - a) # DDIM (eta=0) 的一步,含 CFG 与引导区间 e_un = fn(X, t, LOGW_UNCOND) e_c = fn(X, t, LOGW_COND) w_eff = w if (lo <= t / T <= hi) else 1.0 e_g = e_un + w_eff * (e_c - e_un) x0 = (X - np.sqrt(1.0 - a) * e_g) / np.sqrt(a) X = np.sqrt(a_prev) * x0 + np.sqrt(max(1.0 - a_prev, 0.0)) * e_g 五行里三个细节值得说: w_eff 那一行的 t / T 是归一化噪声档位,区间外直接退回 $w=1$(不是 $w=0$)。这就是 guidance interval 的全部实现。 e_g 可以写成 w_eff * e_c + (1 - w_eff) * e_un,两者数值等价但后者在 $w$ 很大时更容易看出负系数——建议保留后者以免误读。 DDIM 的更新里 $x_0$ 和 $e_g$ 用的是同一个 $e_g$,所以缩放 $e_g$ 会同时改变 $x_0$ 和噪声项,不是单向的。这一点在第 06 节 rescale 那一段会咬人。 先验证去噪器本身是对的。做法是数值积分对拍:固定 20 万真实样本,对每个探测点 $x_t$ 按 $q(x_t|x_0)$ 加权求样本平均,跟闭式解比。 [Q1] 精确去噪器 E[x0|xt] 的数值积分对拍 t alpha_bar 闭式解 E[x0] 数值积分 E[x0] 最大绝对差 1 9.999000e-01 [-0.80383,-0.00497] [-0.80352,-0.00542] 3.00e-02 50 9.710157e-01 [+0.33377,-0.32706] [+0.33434,-0.32566] 1.13e-02 200 6.590385e-01 [-1.63345,-1.01052] [-1.62640,-1.00988] 1.32e-02 500 7.858724e-02 [-0.03520,+0.23795] [-0.03229,+0.24075] 1.03e-02 900 2.752059e-04 [+0.01798,-0.00712] [+0.02195,-0.00467] 4.12e-03 1000 4.035830e-05 [+0.00501,-0.00579] [+0.00899,-0.00330] 4.03e-03 差值在 $10^{-2}$ 量级且随 $t$ 增大而减小:$t$ 大时 $q(x_t|x_0)$ 的权重平、有效样本多,$t$ 小时权重集中、积分噪声大。这是纯蒙特卡洛噪声的水平,闭式解可以放心用。 再验证采样器步数够不够($w=7.5$): steps 软纯度 多样性比 典型度 类心偏移 50 0.9640 0.7490 -2.2718 1.7794 100 0.9637 0.7782 -2.2930 1.7803 200 0.9636 0.7933 -2.3042 1.7809 400 0.9635 0.8024 -2.3112 1.7813 1000 0.9635 0.8055 -2.3136 1.7815 软纯度和类心偏移在 50 步就收敛了,但多样性比和典型度还在爬(0.7490→0.8055,−2.2718→−2.3136)。200 步的典型度距 1000 步约 0.4%,多样性比仍相差约 1.5%,后面全部实验统一用 200 步。这个坑值得记:只看 FID/纯度会以为早就收敛了,多样性类的指标还在漂。 05. 工业级实现对照 参考实现: huggingface/diffusers/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py → StableDiffusionPipeline.__call__ 以 2026-09 的实现为准,去噪循环里的核心是这几行: # expand the latents if we are doing classifier free guidance latent_model_input = torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents if hasattr(self.scheduler, "scale_model_input"): latent_model_input = self.scheduler.scale_model_input(latent_model_input, t) noise_pred = self.unet(latent_model_input, t, encoder_hidden_states=prompt_embeds, ...)[0] if self.do_classifier_free_guidance: noise_pred_uncond, noise_pred_text = noise_pred.chunk(2) noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond) if self.do_classifier_free_guidance and self.guidance_rescale > 0.0: noise_pred = rescale_noise_cfg(noise_pred, noise_pred_text, guidance_rescale=self.guidance_rescale) latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0] 和我的最小实现比,有四处差异,每一处都有理由: 第一,合批不等于只启动一个 kernel。 torch.cat 让一次 UNet 调用同时处理两分支,但网络仍包含很多算子、kernel 和临时张量。理论算术量约翻倍;实际延迟与吞吐取决于 batch、显存、并行和硬件,不能直接断言延迟只涨 30%~60% 或吞吐必然减半。 第二,do_classifier_free_guidance 是一个属性而不是参数: @property def do_classifier_free_guidance(self): return self._guidance_scale > 1 and self.unet.config.time_cond_proj_dim is None 两个条件都值得注意。guidance_scale > 1 意味着 $w=1$ 时无条件分支根本不会被构造,连 negative_prompt_embeds 都不会编码——「免费的 $w=1$」在代码层面是真的。而 time_cond_proj_dim is None 是更关键的一句:当 UNet 配置里存在 time_cond_proj_dim(把引导强度作为时间条件的投影维度注入)时,diffusers 自动关掉双分支 CFG。这是带 guidance embedding 的 UNet 关闭双分支 CFG 的条件。LCM 与 LCM-LoRA 配置并不完全相同,不能仅凭模型名字判断;应核对 checkpoint 和 pipeline 是否使用引导嵌入或双分支。 第三,rescale_noise_cfg 把缩放做成了插值而不是替换: std_text = noise_pred_text.std(dim=list(range(1, noise_pred_text.ndim)), keepdim=True) std_cfg = noise_cfg.std(dim=list(range(1, noise_cfg.ndim)), keepdim=True) noise_pred_rescaled = noise_cfg * (std_text / std_cfg) noise_cfg = guidance_rescale * noise_pred_rescaled + (1 - guidance_rescale) * noise_cfg 原文(Lin 等人 2305.08891 第 3.4 节)的做法是把引导后预测的标准差拉回条件分支的标准差。这里额外乘了一个 $\phi\in[0,1]$ 做插值,注释写得很直白:完全 rescale 会得到「plain looking」的图——拉回标准差的同时也把引导的锐度拉掉了。0.7 是论文/示例中的一种建议值,不是 SDXL pipeline 的通用默认;当前 diffusers 接口 guidance_rescale 默认是 0.0。另外注意 std 是逐样本、跨所有通道与空间位置求的,不是全局统计量。 第四,负向 prompt 走的是同一个条件通道。 negative_prompt_embeds 与 prompt_embeds 拼在同一个 batch 里,chunk(2) 之后无条件分支就是负向 prompt 的预测。所以「负向 prompt」不是 CFG 之外的一个独立功能,它就是无条件分支的语义化替换:把「什么都不给」换成「明确不要什么」。负向 prompt 的效果还取决于编码器、截断长度与模型训练;公式本身不能证明“写太长就失效”。 06. 代价与边界 6.1 算力账:2× 是定义级的 把 SD-v1-5 的 UNet 在 512×512(潜空间 64×64)、fp16、batch=1 下的卷积逐层列出来算一遍: [L1] stable-diffusion-v1-5 UNet,512x512 图(潜空间 64x64),fp16,batch=1 配置 单步 MAC(下限) 层间张量字节(下限) CFG off 1.898e+11 0.13 GB CFG on (w>1) 3.796e+11 0.25 GB 倍率:MAC 2.0000 | 激活 2.0000 上表卷积的参数量合计 517.0 M(SD1.5 UNet 全量约 859 M) 口径说明:这里只数了「一层卷积 = 一个输入张量 + 一个输出张量」,注意力投影、GroupNorm、SiLU 的中间结果、残差分支的暂存都没有计入,所以 0.13 GB 是逐层输入输出张量字节和的局部账本,不是峰值显存,也不是峰值显存的严格下限(不同层可复用内存)。但本文所有结论只用它的比值,而比值 2.0000 是定义级的(batch 维从 1 变 2,权重一份不变),不依赖口径。 算术量翻倍兑现成多少墙钟?实测一个 64×64×320→320 的 3×3 卷积(SD 第一个 ResBlock 的尺寸),CPU + numpy BLAS 取 5 次最小值: batch=1 : 0.0860 s batch=2 : 0.1841 s 时间比 : 2.1405 CPU 上 gemm 是算术受限的,所以比值贴近 2。GPU 上可能因合批提高利用率而使墙钟比小于 2;单位时间出图数由实测总延迟和 batch 决定,也不必精确减半。 换成端到端:20 步 DDIM 是 20 次前向变 40 次,50 步是 50 次变 100 次。 这张图要看什么:三张子图分别回答「算力翻倍」「误差放大」「幅度膨胀」。左图是对数刻度下的 MAC、激活字节、50 步前向次数,CFG on/off 两组柱子的高度差在任何一项上都是同样的 2×——强调它是乘性的,跟模型大小、步数、分辨率都无关。中图是实测误差放大倍数(蓝)对最坏上界 $2w-1$(红虚线),两条线差一个数量级,但形状都是单调的。右图是引导后预测的标准差比 $\text{std}(e_{\text{guided}})/\text{std}(e_{\text{c}})$ 随 $w$ 的曲线,$w=1$ 时是 1.00,$w=7.5$ 时 1.77,$w=25$ 时 4.21——这条曲线就是「过曝」的量化版本。 6.2 误差放大:实测远小于上界,但方向一致 探测点固定为「真实数据前向扩散到 $t$」,与 $w$ 无关,这样排除轨迹漂移的干扰: $t$ $w{=}1$ $w{=}2$ $w{=}3$ $w{=}5$ $w{=}7.5$ $w{=}15$ 最坏上界(同序) 1000 1.000 1.298 1.839 3.037 4.580 9.265 1 / 3 / 5 / 9 / 14 / 29 700 1.000 1.301 1.842 3.042 4.589 9.283 同上 400 1.000 1.389 1.952 3.210 4.834 9.764 同上 200 1.000 1.540 2.184 3.553 5.302 10.596 同上 50 1.000 1.598 2.267 3.663 5.439 10.806 同上 10 1.000 1.599 2.269 3.664 5.438 10.801 同上 $w=15$ 时最坏上界是 29 倍,实测 9.3~10.8 倍。差距来自 3.3 节那个放缩的两处放水:一是最坏范数界允许经系数符号作用后同向叠加;二是它用两支误差的较大范数统一界定。两支预测相关系数 0.9928 不是两支误差相关系数,不能由它证明误差抵消;还要注意表格按条件分支误差归一,而 $(2w-1)$ 界按两支较大误差归一;只有分母一致时才能逐项比较。但随 $w$ 单调放大这件事是所有 $t$ 上一致的,而且低噪声端($t$ 小)放大得更狠——因为两支预测在 $t$ 小时都趋近于真实噪声,误差结构更接近。 6.3 幅度膨胀与 rescale 的边界 $t=300$ 处,探测点来自真实条件样本: $w$ 0.0 1.0 2.0 3.0 5.0 7.5 10.0 15.0 25.0 $\text{std}$ 比 0.9110 1.0000 1.1009 1.2109 1.4493 1.7690 2.1026 2.7914 4.2052 $w=0$(纯无条件)的比是 0.9110——比 1 还小,因为无条件分支要覆盖全部 4 个分量,它的预测更「平均」。从 $w=1$ 往上单调涨到 4.21。 下面只做二维 toy 的 batch 标准差线性 rescale,在 $w=15$ 下扫系数 $\phi$。真实图像实现按每个样本的 C/H/W 统计;二维向量只有两个坐标,不能把此处整批统计量当成原论文的逐样本实现: $\phi$ 多样性比 典型度 软纯度 0.0 0.6716 −5.5530 0.9905 0.3 0.8274 −6.1435 0.9910 0.5 0.9513 −6.5778 0.9913 0.7 1.0944 −7.0504 0.9914 1.0 1.3521 −7.8447 0.9916 结论是:能拉回多样性,但典型度更差(−5.5530 → −7.8447)。机制在 04 节那个细节里已经埋好了——DDIM 的更新中 $x_0$ 是 $e_{\text{guided}}$ 的仿射函数,缩放 $e_{\text{guided}}$ 的幅度等于把 $x_0$ 往「没去噪干净」的方向拽:多样性回来了,是因为你把样本往噪声里推回去了,不是因为它更对了。所以 rescale 修的是观感(过曝),不是分布。这个结论限定在本设定(精确去噪器 + DDIM $\eta=0$);这不能解释或否定真实图像的逐样本 rescale 效果,后者需要在实际模型上比较。 6.4 引导区间:帕累托前沿在哪 把 $w=7.5$ 只开在归一化噪声档位 $t/T$ 的某个区间内(区间外退回 $w=1$): 区间($t/T$) 软纯度 多样性比 典型度 类心偏移 说明 $[0.0,1.0]$ 0.9636 0.7933 −2.3042 1.7809 全程开 $[0.5,1.0]$ 0.9126 0.9571 −0.8896 0.9365 只在高噪声段(链首) $[0.2,0.8]$ 0.9594 0.8167 −2.0693 1.6609 只在中段 $[0.0,0.5]$ 0.9519 0.7672 −1.1345 1.2251 只在低噪声段(链尾) $[0.0,0.2]$ 0.9042 0.8172 −0.0831 0.4222 只在极低噪声段 $[0.0,0.0]$ 0.7409 0.9876 +0.0030 0.0075 全程不开($w=1$ 基线) 再细扫两族区间,用「纯度增益 / 典型度代价」当性价比: 细扫 A:只开链尾(区间 = [0, hi]) 区间 纯度增益 典型度代价 性价比 类心偏移 多样性比 [0.0,0.1] +0.0829 0.0099 8.345 0.1649 0.8899 [0.0,0.2] +0.1632 0.0861 1.896 0.4222 0.8172 [0.0,0.3] +0.1921 0.3080 0.624 0.6888 0.7890 [0.0,0.4] +0.2040 0.6816 0.299 0.9630 0.7737 [0.0,0.5] +0.2110 1.1375 0.185 1.2251 0.7672 [0.0,0.8] +0.2211 2.1335 0.104 1.7045 0.7862 [0.0,1.0] +0.2226 2.3072 0.097 1.7809 0.7933 细扫 B:只开链首(区间 = [lo, 1.0]) 区间 纯度增益 典型度代价 性价比 类心偏移 多样性比 [0.0,1.0] +0.2226 2.3072 0.097 1.7809 0.7933 [0.3,1.0] +0.2154 2.0350 0.106 1.6224 0.8669 [0.5,1.0] +0.1716 0.8926 0.192 0.9365 0.9571 [0.7,1.0] +0.0674 0.1837 0.367 0.2866 0.9813 [0.9,1.0] +0.0101 0.0209 0.481 0.0438 0.9868 读法:全程开的纯度增益是 +0.2226,典型度代价是 2.3072。只开链尾 10%(区间 $[0,0.1]$)就能拿到 +0.0829,代价只有 0.0099,性价比 8.345——是全程的 86 倍。 反过来看 B 族:只开链首 10%($[0.9,1.0]$)的纯度增益只有 +0.0101,也就是说链首那 10% 的引导几乎买不到任何东西,却占了不小的代价份额(单独开它的性价比 0.481,看着还行,是因为它代价绝对值小;但从全程里减掉它,增益只掉 0.0006)。 这张图要看什么:横轴是典型度代价、纵轴是纯度增益,双对数坐标,虚线是等性价比线。A 族(绿,只开链尾)整条曲线压在 B 族(蓝,只开链首)的左上方——同样的代价,A 族拿到的增益更多。两族共同端点是右上角那个「全程开」的点,性价比 0.097,是全图最差。左下角 $[0,0.1]$ 那个点孤零零地挂在性价比 8.345 的位置上。读这张图的结论不是「照抄区间」,而是「常量 $w$ 全程开的那个点,恰好落在帕累托前沿的最差端」。 需要诚实地区分一下:Kynkäänniemi 等人(2404.07724)在 ImageNet-512 上的结论是「引导在链首有害、在链尾基本没必要、只有中段有用」,而我的解析实验里性价比最高的是链尾。两边不能直接对齐——我是二维高斯混合 + 精确去噪器,纯度增益在链上的分布跟 ImageNet + 训练出来的 UNet 不一样。两边真正一致的结论是:常量 $w$ 全程开不是最优,区间应当作为超参暴露出来。 用的时候请以自己模型的实测为准,别照抄任何一边的具体区间。 6.5 什么时候不该用 已经被蒸馏掉的模型:UNet 有 time_cond_proj_dim(引导强度烘进网络)时,再开双分支 CFG 是纯浪费——diffusers 已经帮你自动关了,手写推理代码时得自己关。 $w$ 已经很大还在往大调:从 $w=10$ 到 $w=25$,软纯度只从 0.9744 涨到 0.9985,典型度从 −3.3963 掉到 −9.7613。这一段是纯亏。 多样性是硬指标的场景(数据增广、多样性评测、素材批量生成):CFG 的多样性比在 $w=7.5$ 时已经掉到 0.7933,且没有平台期。 模型本身很弱时:弱去噪器那一栏 $w=15$ 把方差压到真实值的 0.28%。弱模型 + 大 $w$ = 确定性塌缩。这种情况该修模型,不是调 $w$。 07. 经典论文脉络 Classifier-Free Diffusion Guidance(arXiv:2207.12598)——Ho & Salimans。提出用「随机丢弃条件训练出来的同一个网络」替代外置分类器,把 classifier guidance 的对抗梯度换成两支预测的外推。留下的问题是:它默认 $w$ 是全程常量,且没有量化代价。 Common Diffusion Noise Schedules and Sample Steps are Flawed(arXiv:2305.08891)——Lin 等人。两个独立贡献:训练端的 zero-terminal-SNR 调度(让最后一步真的能走到纯噪声)与推理端的 guidance rescale(3.4 节)。Rescale 直接对着「$w$ 大了会过曝」这个现象下刀,是 rescale 方法的来源;6.3 节数值由本文二维 batch 统计 toy 实测,并非论文原表。 Applying Guidance in a Limited Interval Improves Sample and Distribution Quality in Diffusion Models(arXiv:2404.07724)——Kynkäänniemi 等人(NeurIPS 2024)。指出引导在链首有害、链尾基本没必要、只有中段有用,把引导限制在噪声水平的某个区间内,ImageNet-512 的 FID 从 1.81 降到 1.40,并建议在所有用引导的扩散模型里把区间作为超参暴露出来。 Guiding a Diffusion Model with a Bad Version of Itself(arXiv:2406.02507)——Karras 等人。用「训练不足的自己」当引导的负支,替代空条件分支。动机正好是本文 3.3 节那个负系数:既然误差会被放大,那就让负支的误差方向更有用——欠训练模型保留的是低频结构,引导方向因此更「语义」而不是更「纹理」。 Latent Consistency Models(arXiv:2310.04378)——Luo 等人。把「带引导的反向过程」看成一个增广的概率流 ODE,直接蒸馏它的解。这一步之后 $w$ 变成网络的一个输入,每步只需一次前向——这是「干掉 CFG 那 2×」最彻底的一条路,也是知识树里 step_distillation 那篇要展开的内容。 五篇的演进关系是一条很清楚的线:提出机制 → 修观感 → 修区间 → 修负支 → 把机制整个吸收进权重。 前四篇都在「怎么把 CFG 用得更好」,最后一篇是「怎么不再需要它」。 08. 常见误解 误解一:「CFG 就是在采样 $p(x|c)^w p(x)^{1-w}$。」 score 层面恒等,分布层面不成立。参考统计量是交叉熵 $-E_{X\sim Q}[\log P(X)]$。不同列目标不同,不能横向比较“越小越像”;即使固定 $P$,较低交叉熵也可能来自模式坍缩,需与均值、方差或分布距离联读: | $w$ | 样本来源 | $-\mathbb{E}[\log p(x\|c)]$ | $-\mathbb{E}[\log p_{\text{tilde}}]$ | $-\mathbb{E}[\log p(x)]$ | 方差比 | 类心偏移 | |---|---|---|---|---|---|---| | — | 真实条件样本 | 2.5757 | — | 2.8562 | 1.0000 | 0.0000 | | 1.0 | CFG 生成 | 2.5670 | 2.5670 | 2.8532 | 0.9876 | 0.0075 | | 1.0 | 倾斜目标 | 2.5891 | 2.5891 | 2.8706 | 1.0090 | 0.0096 | | 3.0 | CFG 生成 | 2.7405 | 2.4039 | 3.3536 | 0.8193 | 0.7721 | | 3.0 | 倾斜目标 | 2.4824 | 2.3227 | 3.0070 | 0.8944 | 0.3546 | | 7.5 | CFG 生成 | 4.5048 | 3.6926 | 5.1605 | 0.7933 | 1.7809 | | 7.5 | 倾斜目标 | 2.6350 | 2.2258 | 3.2287 | 0.9077 | 0.6028 | | 15.0 | CFG 生成 | 7.7257 | 6.1882 | 8.4092 | 0.6716 | 3.0681 | | 15.0 | 倾斜目标 | 2.9792 | 2.2782 | 3.6030 | 1.0059 | 0.8828 | $w=1$ 时两行几乎重合(差 0.02 左右),这个说法是对的。$w=3$ 已经开始分叉(类心偏移 0.7721 对 0.3546)。到 $w=7.5$,CFG 生成样本相对倾斜目标的交叉熵是 3.6926,而倾斜目标对自己的交叉熵是 2.2258,差 1.47 nats;类心偏移 1.78 对 0.60,差 3 倍。固定目标下两个期望显著不同,可以否定两分布相等;相近则不能证明分布相同。$w=1$ 在精确模型下成立,其他 $w$ 不能仅凭“很小”保证等价。归一化代码已保留 log-sum-exp 的偏移量,$w=1$ 时 logp_tilt_norm == logp_cond 可直接检验。 误解二:「$w$ 越大越准。」 纯度上去的同时典型度一路往下,而且没有平台期:$w=15$ 时典型度 −5.5530,$w=25$ 时 −9.7613。它变「准」的那一维是「属于目标类的程度」,代价是「属于真实数据的程度」。 误解三:「CFG 不增加参数,所以是免费的。」 MAC 从 $1.898\times10^{11}$ 到 $3.796\times10^{11}$,比值 2.0000;50 步 DDIM 的前向次数从 50 到 100。这是推理成本里最容易被漏掉的一项。 误解四:「把 $w$ 调小一点就省算力。」 不省。只要 $w>1$,do_classifier_free_guidance 就是 True,两个分支都要跑。$w=1.01$ 和 $w=30$ 的算力完全一样。 误解五:「rescale 能修 CFG 的过曝。」 在精确去噪器下它把多样性拉回来了(0.6716→1.3521)但典型度更差(−5.5530→−7.8447)。它修的是幅度观感,不是分布。把 rescale 当成「可以放心加大 $w$ 的许可证」是错的。 误解六:「$w<1$ 就是减弱引导。」 只有 $0<w<1$ 时两个系数都为正,预测是在两支之间插值,输出一般不是条件/无条件分布的概率混合而不是「弱一点的条件采样」。实测 $w=0$ 时多样性比是 1.2512(比真实条件还宽)、类心偏移 0.6651(往相反方向偏)。这是一个不同的分布,不是同一个分布的弱化版。 09. 动手验证 三个小实验,脚本在附录里,全部只依赖 numpy 与 matplotlib。 实验一:确认你的去噪器有没有误差。 跑 cfg_gmm.py 的 Q1 段,看闭式解与数值积分的差。差在 $10^{-2}$ 量级且随 $t$ 增大而减小,说明是蒙特卡洛噪声;如果差随 $w$ 变,说明你把 CFG 写进了去噪器本身。 实验二:扫你自己的 $w$。 改 cfg_gmm.py 里的 SPECS(换成你关心的类心距离与方差结构),跑 Q2。本文那张表的形状是:软纯度在 $w\approx3$ 前快速上升、之后变平;典型度与类心偏移全程线性恶化;多样性比是唯一一个在中间段有一点非单调的指标(0.8042 → 0.7933 → 0.7603,这三点实际单调下降)。如果你扫出来的曲线在这三项上形状一致,说明机制对上了。 实验三:量你自己的引导区间。 跑 cfg_gmm.py 的 Q5 段(细扫 A 与细扫 B),画出「纯度增益 vs 典型度代价」的双对数图。看两件事:你的曲线是不是也在全程开那个点性价比最低;以及 A 族(链尾)与 B 族(链首)哪一条压在左上。这张图应当成为你调 $w$ 之前先看的图,因为它告诉你性价比最高的区间在哪,而不是告诉你 $w$ 该取几。 实验四:算你自己模型的账。 把 cfg_lab.py 里的 CONVS 换成你的 UNet 配置(stage、分辨率、输入输出通道、重复次数),跑一遍看 MAC 与激活的比值是不是 2.0000。该账本比值由样本维决定;分两次调用不改变理论 MAC。实测峰值显存或延迟不等于 2 并不能说明实现错误。 10. 延伸阅读 读这篇之前建议先看: DDPM 训练目标与采样流程——$\varepsilon$ 预测器、$L_{\text{simple}}$、以及采样循环怎么写。本文 04 节的 DDIM 一步直接用了那篇的结论。 扩散过程的前向与反向推导——$\bar\alpha_t$、后验方差、以及 $\varepsilon$ 与 score 的换算关系(3.2 节那一步在那篇推过)。 从 DDIM 到高阶采样器——本文全部实验用 DDIM $\eta=0$,换采样器会改变引导误差累积的方式。 VAE 的 ELBO 怎么拆——条件生成的「条件」到底以什么形式进入模型,这是 CFG 能成立的前提。 读完之后: 少步蒸馏:从 50 步到 4 步——沿着 07 节最后一篇继续,看 $w$ 是怎么被烘进网络、从而把 2× 变回 1× 的。(写作中) 回到知识树:本文是「生成范式」方向的一级节点,往上接 DDPM,往下接少步蒸馏。 附录:完整代码 09 节用到的脚本全文如下(cfg_gmm.py、cfg_lab.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 cfg_gmm.py # -*- coding: utf-8 -*- """cfg_gmm.py —— 在一个「去噪器可解析求解」的高斯混合模型上把 CFG 做实。 为什么非要用 GMM: CFG 需要一支条件去噪器和一支无条件去噪器。用真网络的话两支都有训练误差, 观察到的任何形变都分不清是 CFG 造成的还是训歪了。而 GMM 的 MMSE 去噪器 E[x0|xt] 有闭式解,两支都是**精确**的,于是生成分布的形变只能来自 CFG 本身。 另外 p_t(x_t|c) 也有解析式,可以直接验证「CFG 是不是在从倾斜分布 p(x|c)^w p(x)^{1-w} 里采样」这个命题。 本脚本回答六个问题(全部是实跑数字,不是推测): Q1 精确去噪器对不对?(数值积分对拍) Q2 w 扫描:条件纯度 / 多样性 / 典型度 怎么变? Q3 CFG 采出来的分布,是不是那个倾斜分布? Q4 w 把去噪器的**误差**放大了多少倍?(最坏上界 2w-1) Q5 只在某个噪声区间开引导,比全程开好吗? Q6 引导后的预测幅度膨胀多少?rescale 能不能压回去? 运行: /usr/local/bin/python3 cfg_gmm.py 依赖: numpy(无 torch,全部闭式解 + 向量化) """ import numpy as np # ═════════════════════════════ 1. 数据:两组、四个高斯分量 ═════════════════════════════ # (均值, 各向同性方差, 无条件权重, 类别标签) # 设计意图:两类故意重叠(类心相距约 1.3,分量标准差 0.6~1.0), # 这样 w=1 时条件采样仍会「漏」到对面,给 CFG 留出可观测的改善空间。 # 两类各有一个「紧」分量和一个「松」分量 —— 方差不齐是后面「模式漂进尾巴」的关键。 SPECS = [ ((-0.90, -0.05), 0.36, 0.25, 0), ((-0.10, -0.85), 1.00, 0.25, 0), ((0.90, 0.05), 0.36, 0.25, 1), ((0.10, 0.85), 1.00, 0.25, 1), ] MU = np.array([s[0] for s in SPECS], dtype=float) # [K, 2] VAR = np.array([s[1] for s in SPECS], dtype=float) # [K] PI = np.array([s[2] for s in SPECS], dtype=float) # [K] CLS = np.array([s[3] for s in SPECS], dtype=int) # [K] K = len(SPECS) DIM = 2 TARGET_CLS = 1 # 本文统一用「条件 = 类别 1」做演示 # ── 调度:DDPM 线性 beta,T=1000(与 ddpm 那篇同一套口径)── T = 1000 BETAS = np.linspace(1e-4, 0.02, T) ALPHAS = 1.0 - BETAS ABAR = np.cumprod(ALPHAS) # alpha_bar,下标 0 对应 t=1 def abar(t): """alpha_bar_t,t 从 1 开始(与论文记号一致)。""" return ABAR[t - 1] def safe_log(w): """对数,0 分量记为 -inf 而不报警告。""" w = np.asarray(w, dtype=float) out = np.full_like(w, -np.inf) m = w > 0 out[m] = np.log(w[m]) return out # ═════════════════════════════ 2. 精确去噪器 ═════════════════════════════ def logpdf_comps(X, a_bar): """每个分量对 x_t 的对数密度:x_t|k ~ N(sqrt(a) mu_k, a var_k + (1-a))。""" s = np.sqrt(a_bar) m = s * MU # [K, 2] v = a_bar * VAR + (1.0 - a_bar) # [K] d2 = ((X[:, None, :] - m[None, :, :]) ** 2).sum(-1) # [N, K] return -0.5 * (d2 / v[None, :] + DIM * np.log(2 * np.pi * v)[None, :]) def posterior(X, a_bar, logw): """分量后验权重 r_k = p(k | x_t)。logw 是分量的对数先验。""" lp = logpdf_comps(X, a_bar) + logw[None, :] lp -= lp.max(axis=1, keepdims=True) r = np.exp(lp) return r / r.sum(axis=1, keepdims=True) def x0_hat(X, a_bar, logw): """MMSE 估计 E[x0 | x_t](精确闭式解)。 分量内部是高斯的,所以 E[x0|x_t,k] 有闭式解: mu_k + S_k sqrt(a) (a S_k + (1-a) I)^{-1} (x_t - sqrt(a) mu_k) 再按后验权重 r_k 加权。S_k = var_k * I,所以增益是个标量。 """ s = np.sqrt(a_bar) v = a_bar * VAR + (1.0 - a_bar) gain = VAR * s / v # [K] r = posterior(X, a_bar, logw) # [N, K] per = MU[None, :, :] + gain[None, :, None] * (X[:, None, :] - s * MU[None, :, :]) return (r[:, :, None] * per).sum(axis=1) # [N, 2] def eps_hat(X, t, logw): """epsilon 预测:eps = (x_t - sqrt(a) x0_hat) / sqrt(1-a)。""" a = abar(t) return (X - np.sqrt(a) * x0_hat(X, a, logw)) / np.sqrt(1.0 - a) def cond_logw(cls): """条件分支的分量对数先验:只保留该类内部分量,类内权重重新归一化。""" w = np.where(CLS == cls, PI, 0.0) w = w / w.sum() return safe_log(w) LOGW_UNCOND = safe_log(PI) # 无条件:全混合 LOGW_COND = cond_logw(TARGET_CLS) # 条件:类别 TARGET_CLS # ── 弱去噪器:把整簇拟合成「一个高斯」之后的线性维纳滤波 ── # 给 Q4 用的「不完美模型」:它是该高斯下的最优线性去噪器,但真实数据是混合体, # 所以误差非零、且随 t 变化 —— 正好用来看 w 的放大倍数。 def single_gauss(logw): """按分量权重算出混合体的均值与平均方差。""" w = np.exp(logw) w = w / w.sum() mu = (w[:, None] * MU).sum(0) var = (w * (VAR + (MU ** 2).sum(1))).sum() - (mu ** 2).sum() return mu, var / DIM def weak_eps_hat(X, t, logw): """单高斯近似的线性去噪器(有误差)。""" a = abar(t) mu, var = single_gauss(logw) s = np.sqrt(a) v = a * var + (1.0 - a) gain = var * s / v x0 = mu[None, :] + gain * (X - s * mu[None, :]) return (X - s * x0) / np.sqrt(1.0 - a) # ═════════════════════════════ 3. 采样器(DDIM, eta=0) ═════════════════════════════ def ddim_timesteps(steps): """均匀步长的 DDIM 时间步序列,从 T 递减到 1。""" stride = max(T // steps, 1) return list(range(T, 0, -stride)) def ddim_sample(x_T, steps=200, w=1.0, interval=(0.0, 1.0), weak=False, rescale=0.0): """确定性 DDIM 采样(eta=0),带 CFG。 x_T : [N, 2] 纯噪声起点 w : 引导强度(w=1 即纯条件,w=0 即纯无条件) interval : 引导生效的 t/T 区间 (lo, hi)。区间外退化为 w=1。 weak : True 则用单高斯弱去噪器(Q4/Q6 用) rescale : 二维玩具的 batch 标准差线性 rescale 系数 phi;不等于图像逐样本 rescale """ fn = weak_eps_hat if weak else eps_hat ts = ddim_timesteps(steps) X = x_T.copy() lo, hi = interval for i, t in enumerate(ts): a = abar(t) a_prev = abar(ts[i + 1]) if i + 1 < len(ts) else 1.0 e_un = fn(X, t, LOGW_UNCOND) e_c = fn(X, t, LOGW_COND) frac = t / T w_eff = w if (lo <= frac <= hi) else 1.0 e_g = e_un + w_eff * (e_c - e_un) # = w_eff*e_c + (1-w_eff)*e_un if rescale > 0.0: # 二维 toy 用整批统计量;图像实现应按样本跨 C/H/W 统计。 # 沿用线性插值,而不是把标准差比取 phi 次幂。 s_g, s_c = e_g.std(), e_c.std() if s_g > 1e-12: e_g = (1.0 - rescale) * e_g + rescale * e_g * (s_c / s_g) x0 = (X - np.sqrt(1.0 - a) * e_g) / np.sqrt(a) X = np.sqrt(a_prev) * x0 + np.sqrt(max(1.0 - a_prev, 0.0)) * e_g return X # ═════════════════════════════ 4. 度量 ═════════════════════════════ def logsumexp_rows(lp): m = lp.max(axis=1, keepdims=True) return m[:, 0] + np.log(np.exp(lp - m).sum(axis=1)) def logp_data(X): """真实数据分布(无条件 GMM)的对数密度。""" return logsumexp_rows(logpdf_comps(X, 1.0) + safe_log(PI)[None, :]) def logp_cond(X): """真实条件分布 p(x|c=TARGET_CLS) 的对数密度。""" return logsumexp_rows(logpdf_comps(X, 1.0) + LOGW_COND[None, :]) def class_posterior(X): """p(类别=TARGET_CLS | x) —— 用真实 GMM 算,作为「软纯度」。""" lp = logpdf_comps(X, 1.0) + safe_log(PI)[None, :] m = lp.max(axis=1, keepdims=True) r = np.exp(lp - m) r /= r.sum(axis=1, keepdims=True) return (r * (CLS[None, :] == TARGET_CLS)).sum(axis=1) def spread(X): """分布的「宽度」:每维方差的平均。""" return float(np.mean(X.var(axis=0))) def summarize(X, Xref, tag=""): """一组样本的核心指标。Xref 是真实条件分布 p(x|c) 的样本,作为基准。""" d2 = ((X[:, None, :] - MU[None, :, :]) ** 2).sum(-1) / VAR[None, :] return dict( tag=tag, purity=float(class_posterior(X).mean()), hard=float((CLS[np.argmin(d2, axis=1)] == TARGET_CLS).mean()), div=float(spread(X) / spread(Xref)), typicality=float(logp_data(X).mean() - logp_data(Xref).mean()), bias=float(np.linalg.norm(X.mean(0) - Xref.mean(0))), ) def sample_data(rng, n, cls=None): """从真实 GMM 采 n 个样本;cls 非空则只采该类的样本。""" w = PI.copy() if cls is not None: w = np.where(CLS == cls, PI, 0.0) w = w / w.sum() k = rng.choice(K, size=n, p=w) return MU[k] + np.sqrt(VAR[k])[:, None] * rng.standard_normal((n, DIM)) # ═══════════════════════ 5. 倾斜分布 p(x|c)^w p(x)^{1-w} ═══════════════════════ def tilted_logpdf(X, w): """未归一化的倾斜对数密度:w*log p(x|c) + (1-w)*log p(x)。""" return w * logp_cond(X) + (1.0 - w) * logp_data(X) def _tilt_grid(w, lo=-9.0, hi=9.0, n=420): """在方格上算归一化后的倾斜密度,返回 (xs, P[ny, nx]),P 求和为 1。""" xs = np.linspace(lo, hi, n) GX, GY = np.meshgrid(xs, xs) P = np.column_stack([GX.ravel(), GY.ravel()]) lp = tilted_logpdf(P, w) lp -= lp.max() d = np.exp(lp).reshape(n, n) return xs, d / d.sum() _TILT_CACHE = {} def logp_tilt_norm(X, w, lo=-9.0, hi=9.0, n=420): """归一化后的 log p_tilde(x)(归一化常数由方格数值积分得到)。""" key = (float(w), float(lo), float(hi), int(n)) if key not in _TILT_CACHE: xs = np.linspace(lo, hi, n) GX, GY = np.meshgrid(xs, xs) lp = tilted_logpdf(np.column_stack([GX.ravel(), GY.ravel()]), w) offset = lp.max() d = np.exp(lp - offset).reshape(n, n) Z = d.sum() * (xs[1] - xs[0]) ** 2 _TILT_CACHE[key] = offset + np.log(Z) return tilted_logpdf(X, w) - _TILT_CACHE[key] def sample_tilted(rng, w, n): """从倾斜分布 p(x|c)^w p(x)^{1-w} 采样(方格离散近似 + 格内抖动)。""" xs, d = _tilt_grid(w) flat = d.ravel() idx = rng.choice(len(flat), size=n, p=flat / flat.sum()) iy, ix = np.unravel_index(idx, d.shape) step = xs[1] - xs[0] return np.column_stack([xs[ix] + (rng.random(n) - 0.5) * step, xs[iy] + (rng.random(n) - 0.5) * step]) # ═════════════════════════════ 6. 主流程 ═════════════════════════════ def base_rng(): """主实验与画图共用的随机流起点。 两边必须 draw 相同次数、相同顺序,否则图上的数字和正文表格会对不上。 """ rng = np.random.default_rng(20260928) Xref = sample_data(rng, 40000, cls=TARGET_CLS) return rng, Xref def main(): rng, Xref = base_rng() N = 20000 STEPS = 200 print("=" * 78) print("CFG 实跑账本 —— 数据:2D 高斯混合,4 分量 / 2 类,条件 = 类别 %d" % TARGET_CLS) print("=" * 78) # ── Q1:精确去噪器对拍(数值积分) ── print("\n[Q1] 精确去噪器 E[x0|xt] 的数值积分对拍") print(" 做法:固定 20 万真实样本,对每个探测点 xt 按 q(xt|x0) 加权求样本平均。") print(" t alpha_bar 闭式解 E[x0] 数值积分 E[x0] 最大绝对差") mc_rng = np.random.default_rng(7) X0 = sample_data(mc_rng, 200000) for t in (1, 50, 200, 500, 900, 1000): a = abar(t) Xt = np.sqrt(a) * X0 + np.sqrt(1 - a) * mc_rng.standard_normal(X0.shape) idx = mc_rng.integers(0, len(X0), size=60) xt_probe = Xt[idx] mc = np.empty_like(xt_probe) s = np.sqrt(a) for j, xp in enumerate(xt_probe): # 逐点算,避免大临时矩阵 # 权重就是 q(xt|x0) ∝ exp(-||xt - sqrt(a) x0||^2 / (2(1-a))) d2 = ((xp[None, :] - s * X0) ** 2).sum(-1) ww = np.exp(-0.5 * (d2 - d2.min()) / (1 - a)) # 减最小值防下溢 ww /= ww.sum() mc[j] = (X0 * ww[:, None]).sum(0) cf = x0_hat(xt_probe, a, LOGW_UNCOND) # 探测数据来自全混合,故用无条件先验 print(" %4d %.6e [%+.5f,%+.5f] [%+.5f,%+.5f] %.2e" % (t, a, cf[0, 0], cf[0, 1], mc[0, 0], mc[0, 1], np.abs(mc - cf).max())) print(" → 差值在 1e-2 量级;t 大时权重平、有效样本多,t 小时权重集中、积分噪声大。") # 基准 print("\n基准:真实条件分布 p(x|c=%d)" % TARGET_CLS) print(" 每维方差 %.4f | 均值 [%+.4f,%+.4f] | 平均 log p_data %.4f | 软纯度 %.4f" % (spread(Xref), Xref.mean(0)[0], Xref.mean(0)[1], logp_data(Xref).mean(), class_posterior(Xref).mean())) x_T0 = rng.standard_normal((N, DIM)) # ── 步长收敛检查 ── print("\n[Q0] DDIM 步数收敛检查(w=7.5,看 200 步够不够;1000 步 = 全量 T)") print(" steps 软纯度 多样性比 典型度 类心偏移") for st in (50, 100, 200, 400, 1000): X = ddim_sample(x_T0, steps=st, w=7.5) m = summarize(X, Xref) print(" %4d %.4f %.4f %+.4f %.4f" % (st, m["purity"], m["div"], m["typicality"], m["bias"])) # ── Q2:w 扫描 ── print("\n[Q2] 引导强度 w 扫描(精确去噪器,DDIM %d 步,N=%d)" % (STEPS, N)) print(" w 软纯度 硬纯度 多样性比 典型度(nats) 类心偏移") r0 = summarize(Xref, Xref, tag="真实条件") print(" 真实条件样本(参照) %.4f %.4f %.4f %+.4f %.4f" % (r0["purity"], r0["hard"], r0["div"], r0["typicality"], r0["bias"])) sweep = {} for w in (0.0, 1.0, 2.0, 3.0, 5.0, 7.5, 10.0, 15.0, 25.0): X = ddim_sample(x_T0, steps=STEPS, w=w) m = summarize(X, Xref) sweep[w] = m print(" %5.1f %.4f %.4f %.4f %+.4f %.4f" % (w, m["purity"], m["hard"], m["div"], m["typicality"], m["bias"])) print(" 软纯度 = E[p(c|x)](真实后验);硬纯度 = 马氏最近分量属于目标类的比例;") print(" 多样性比 = 生成样本每维方差 / 真实条件方差;") print(" 典型度 = 生成样本与真实样本的平均 log p_data 之差(负 = 平均落在较低密度区,不代表支撑之外)。") # ── Q3:CFG 采的是不是倾斜分布 ── print("\n[Q3] CFG 采出的分布 vs 倾斜目标 p(x|c)^w p(x)^{1-w}") print(" 统计量:固定目标的交叉熵差可否定同分布;更小不保证更像目标") print(" w 样本来源 -E[log p(x|c)] -E[log p_tilde] -E[log p(x)] 方差比 类心偏移") tilt_rows = {} # 参照行:真实条件分布自己的交叉熵(作为「完美拟合」的刻度) print(" ---- 真实条件样本 %10.4f %10s %10.4f %6.4f %6.4f" % (-float(logp_cond(Xref).mean()), "(w=1 时同)", -float(logp_data(Xref).mean()), spread(Xref) / spread(Xref), float(np.linalg.norm(Xref.mean(0) - Xref.mean(0))))) for w in (1.0, 3.0, 7.5, 15.0): Xg = ddim_sample(x_T0, steps=STEPS, w=w) Xt = sample_tilted(rng, w, 20000) for name, X in (("CFG 生成", Xg), ("倾斜目标", Xt)): row = dict( ce_cond=-float(logp_cond(X).mean()), ce_tilt=-float(logp_tilt_norm(X, w).mean()), ce_data=-float(logp_data(X).mean()), div=float(spread(X) / spread(Xref)), bias=float(np.linalg.norm(X.mean(0) - Xref.mean(0))), ) tilt_rows[(w, name)] = row print(" %4.1f %-11s %10.4f %10.4f %10.4f %6.4f %6.4f" % (w, name, row["ce_cond"], row["ce_tilt"], row["ce_data"], row["div"], row["bias"])) print(" ↑ 结合均值和方差判断;仅交叉熵接近不能证明同分布") # ── Q4:误差放大(固定探测分布) ── print("\n[Q4] 弱去噪器(单高斯近似)下,w 把误差放大了多少倍") print(" 探测点固定为「真实数据前向扩散到 t」,与 w 无关,排除轨迹漂移的干扰。") ws4 = (1.0, 2.0, 3.0, 5.0, 7.5, 15.0) print(" t " + " ".join("w=%-4g" % w for w in ws4) + " 最坏上界 2w-1(同序)") probe = sample_data(rng, 20000) amp_by_t = {} for t in (1000, 700, 400, 200, 50, 10): a = abar(t) Xt = np.sqrt(a) * probe + np.sqrt(1 - a) * rng.standard_normal(probe.shape) ec = weak_eps_hat(Xt, t, LOGW_COND) - eps_hat(Xt, t, LOGW_COND) eu = weak_eps_hat(Xt, t, LOGW_UNCOND) - eps_hat(Xt, t, LOGW_UNCOND) base = np.linalg.norm(ec, axis=1).mean() row = [] for w in ws4: amp = np.linalg.norm(w * ec + (1.0 - w) * eu, axis=1).mean() row.append(amp / max(base, 1e-12)) amp_by_t[t] = row print(" %4d " % t + " ".join("%6.3f" % v for v in row) + " " + " ".join("%.0f" % (2 * w - 1) for w in ws4)) print(" → 实测远小于最坏上界:两支误差高度相关、互相抵消;但随 w 单调放大是一致的。") print("\n 弱去噪器下端到端效果(w 越大塌得越狠):") print(" w 生成方差/真条件 软纯度") for w in ws4: X = ddim_sample(x_T0, steps=STEPS, w=w, weak=True) m = summarize(X, Xref) print(" %5.1f %.4f %.4f" % (w, m["div"], m["purity"])) # ── Q5:引导区间 ── print("\n[Q5] 只在某个噪声区间开引导(w=7.5,精确去噪器,DDIM %d 步)" % STEPS) print(" 区间(t/T) 软纯度 多样性比 典型度 类心偏移 说明") intervals = [ ((0.0, 1.0), "全程开"), ((0.5, 1.0), "只在高噪声段(链的前半)"), ((0.2, 0.8), "只在中段"), ((0.0, 0.5), "只在低噪声段(链的后半)"), ((0.0, 0.2), "只在极低噪声段"), ((0.0, 0.0), "全程不开(w=1 基线)"), ] interval_rows = {} base = None for (lo, hi), name in intervals: if (lo, hi) == (0.0, 0.0): X = ddim_sample(x_T0, steps=STEPS, w=1.0) else: X = ddim_sample(x_T0, steps=STEPS, w=7.5, interval=(lo, hi)) m = summarize(X, Xref) interval_rows[(lo, hi)] = dict(name=name, **m) if (lo, hi) == (0.0, 0.0): base = m print(" [%.1f,%.1f] %.4f %.4f %+.4f %.4f %s" % (lo, hi, m["purity"], m["div"], m["typicality"], m["bias"], name)) # 细扫:把「区间末端从哪切」和「区间起点从哪切」分别扫一遍,看帕累托前沿 print("\n 细扫 A:只看链的后半段(区间 = [0, hi]),hi 从 0.1 扫到 1.0") print(" 区间 纯度增益 典型度代价 性价比 类心偏移 多样性比") fine_a = {} for hi in (0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.8, 1.0): X = ddim_sample(x_T0, steps=STEPS, w=7.5, interval=(0.0, hi)) m = summarize(X, Xref) g = m["purity"] - base["purity"] c = -(m["typicality"] - base["typicality"]) fine_a[hi] = dict(gain=g, cost=c, ratio=g / c, **m) print(" [0.0,%.1f] %+.4f %.4f %6.3f %.4f %.4f" % (hi, g, c, g / c, m["bias"], m["div"])) print("\n 细扫 B:只看链的前半段(区间 = [lo, 1.0]),lo 从 0.0 扫到 0.9") print(" 区间 纯度增益 典型度代价 性价比 类心偏移 多样性比") fine_b = {} for lo in (0.0, 0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9): X = ddim_sample(x_T0, steps=STEPS, w=7.5, interval=(lo, 1.0)) m = summarize(X, Xref) g = m["purity"] - base["purity"] c = -(m["typicality"] - base["typicality"]) fine_b[lo] = dict(gain=g, cost=c, ratio=g / c, **m) print(" [%.1f,1.0] %+.4f %.4f %6.3f %.4f %.4f" % (lo, g, c, g / c, m["bias"], m["div"])) # ── Q6:幅度膨胀 + rescale ── print("\n[Q6] 引导后预测的幅度膨胀(t=300,探测点来自真实条件样本)") Xp_base = sample_data(rng, 8000, cls=TARGET_CLS) a300 = abar(300) Xp = np.sqrt(a300) * Xp_base + np.sqrt(1 - a300) * rng.standard_normal(Xp_base.shape) e_c = eps_hat(Xp, 300, LOGW_COND) e_u = eps_hat(Xp, 300, LOGW_UNCOND) rho = float(np.corrcoef(e_c.ravel(), e_u.ravel())[0, 1]) print(" 两分支预测的相关系数 rho = %.4f(高度相关 -> 膨胀来自「差值方向」)" % rho) print(" w std(eps_guided)/std(eps_cond)") infl = {} for w in (0.0, 1.0, 2.0, 3.0, 5.0, 7.5, 10.0, 15.0, 25.0): e_g = e_u + w * (e_c - e_u) infl[w] = float(e_g.std() / e_c.std()) print(" %5.1f %.4f" % (w, infl[w])) print("\n guidance rescale 的效果(w=15):") print(" phi 多样性比 典型度 软纯度") for phi in (0.0, 0.3, 0.5, 0.7, 1.0): X = ddim_sample(x_T0, steps=STEPS, w=15.0, rescale=phi) m = summarize(X, Xref) print(" %.1f %.4f %+.4f %.4f" % (phi, m["div"], m["typicality"], m["purity"])) print(" → 本设定(精确去噪器 + DDIM eta=0)下 rescale 把多样性拉回来了,") print(" 但典型度更差:它缩放的是 eps 的幅度,而 x0 是 eps 的仿射函数,") print(" 缩放幅度等于把 x0 往「没去噪干净」的方向拽。") print("\n" + "=" * 78) print("全部数字由本脚本实跑产生。") print("=" * 78) if __name__ == "__main__": main() cfg_lab.py # -*- coding: utf-8 -*- """cfg_lab.py —— CFG 的代价账本:算术量、激活显存、实测耗时。 这一篇的立论是「CFG 让每步算两遍,是推理成本里最容易被忽视的 2×」。 这句话里「2×」是**定义级精确**的(batch 维翻倍、权重不变), 但读者真正想知道的是:这个 2× 在真机上兑现成多少墙钟时间。 本脚本做三件事: L1 按 stable-diffusion-v1-5 的 UNet 配置手算一份特征图账本(fp16), 给出「单步激活字节」和「单步 MAC」这两个绝对量级。 L2 用 numpy 真跑一个 64x64x320 的 3x3 卷积,实测 batch=1 与 batch=2 的时间比。 (CPU + BLAS,只用于说明算术量翻倍在时间上兑现的程度;GPU 上数字会不同。) L3 把「每步两遍」换算成端到端:50 步采样一共多了多少次前向。 运行: /usr/local/bin/python3 cfg_lab.py 依赖: numpy """ import time import numpy as np # ───────────────────────── L1. UNet 配置与账本 ───────────────────────── # stable-diffusion-v1-5 的 UNet(以 2026-09 时 diffusers 的配置为准): # block_out_channels = [320, 640, 1280, 1280],潜空间 64x64(对应 512x512 图) # 下面只列卷积:ResBlock 内部是 GroupNorm-SiLU-Conv3x3-GroupNorm-SiLU-Conv3x3。 # 注意力模块的 QKV / 输出投影、时间嵌入 MLP 都不在表里 —— 它们参数不少, # 但激活量远小于特征图,且跨注意力在 SD 里只作用于 32/16/8 三个尺度。 CONVS = [ # (阶段名, 空间尺寸, in_ch, out_ch, 重复次数) ("conv_in", 64, 4, 320, 1), ("down0.resnet", 64, 320, 320, 4), # 2 个 ResBlock x 2 个 conv ("down0.downsample", 64, 320, 320, 1), ("down1.resnet", 32, 320, 640, 2), ("down1.resnet", 32, 640, 640, 2), ("down1.downsample", 32, 640, 640, 1), ("down2.resnet", 16, 640, 1280, 2), ("down2.resnet", 16, 1280, 1280, 2), ("down2.downsample", 16, 1280, 1280, 1), ("down3.resnet", 8, 1280, 1280, 4), ("mid.resnet", 8, 1280, 1280, 4), ("up0.resnet", 8, 2560, 1280, 3), ("up0.resnet", 8, 1280, 1280, 3), ("up1.resnet", 16, 2560, 1280, 3), ("up1.resnet", 16, 1280, 1280, 3), ("up2.resnet", 32, 1920, 640, 3), ("up2.resnet", 32, 640, 640, 3), ("up3.resnet", 64, 960, 320, 3), ("up3.resnet", 64, 320, 320, 3), ("conv_out", 64, 320, 4, 1), ] def ledger(dtype_bytes=2, batch=1, with_cfg=False): """算一份特征图账本。 dtype_bytes : fp16 = 2 字节 batch : 一次前向同时处理的样本数 with_cfg : True 则 batch 翻倍(无条件分支拼在 batch 维里) 返回 (总 MAC, 层间张量字节总和, 参数量) """ b = batch * (2 if with_cfg else 1) mac = 0 act = 0 params = 0 for _, hw, cin, cout, rep in CONVS: n = hw * hw mac += b * n * cin * 9 * cout * rep # 一个卷积要留着输入、要写出输出,两块都算层间张量 act += b * n * (cin + cout) * dtype_bytes * rep params += cin * 9 * cout * rep return mac, act, params # ───────────────────────── L2. 实测:一个 3x3 卷积,batch 1 vs 2 ───────────────────────── def conv3x3(X, W): """X: [B, C, H, W],W: [Co, C, 3, 3] —— im2col + 一次大矩阵乘。""" B, C, H, Wd = X.shape Co = W.shape[0] Xp = np.pad(X, ((0, 0), (0, 0), (1, 1), (1, 1))) win = np.lib.stride_tricks.sliding_window_view(Xp, (3, 3), axis=(2, 3)) cols = win.reshape(B, C * 9, H * Wd).transpose(0, 2, 1).reshape(B * H * Wd, C * 9) out = cols @ W.reshape(Co, C * 9).T return out.reshape(B, H, Wd, Co).transpose(0, 3, 1, 2) def time_conv(hw=64, cin=320, cout=320, batch=1, repeat=5): """对一个具体尺寸的 3x3 卷积计时,取 repeat 次的最小值。""" rng = np.random.default_rng(11) X = rng.standard_normal((batch, cin, hw, hw), dtype=np.float32) W = (rng.standard_normal((cout, cin, 3, 3), dtype=np.float32) / np.sqrt(cin * 9)) conv3x3(X, W) # 预热 best = float("inf") for _ in range(repeat): t0 = time.perf_counter() conv3x3(X, W) best = min(best, time.perf_counter() - t0) return best def main(): print("=" * 78) print("CFG 的代价账本") print("=" * 78) print("\n[L1] stable-diffusion-v1-5 UNet,512x512 图(潜空间 64x64),fp16,batch=1") print(" (只算上表列出的卷积;注意力矩阵与 autograd 临时缓冲不计)") print(" ──────────────────────────────────────────────────────────────") print(" %-14s %18s %18s" % ("配置", "单步 MAC(下限)", "层间张量(下限)")) for tag, cfg in (("CFG off", False), ("CFG on (w>1)", True)): mac, act, par = ledger(with_cfg=cfg) print(" %-14s %18.3e %18s" % (tag, mac, "%.3f GB" % (act / 1024 ** 3))) mac0, act0, par0 = ledger(with_cfg=False) mac1, act1, _ = ledger(with_cfg=True) print(" ──────────────────────────────────────────────────────────────") print(" 倍率:MAC %.4f | 激活 %.4f ← 这两个 2 是定义级精确的" % (mac1 / mac0, act1 / act0)) print(" 「下限」口径说明:上表只列了卷积,每个卷积只算 1 份输入 + 1 份输出。") print(" 注意力投影、GroupNorm/SiLU 的中间张量、残差分支都没计进去,") print(" 真实峰值比这两个数大。本文只用它给量级,结论只依赖「翻倍」这个比值。") print(" 上表卷积的参数量合计 %.1f M(SD1.5 UNet 全量约 859 M," "差额是注意力投影与时间嵌入 MLP)" % (par0 / 1e6)) print("\n[L2] 实测:一个 64x64x320 -> 320 的 3x3 卷积(SD 第一个 ResBlock 的尺寸)") print(" CPU + numpy BLAS,取 5 次最小值。用来看算术量翻倍兑现成多少墙钟。") t1 = time_conv(batch=1) t2 = time_conv(batch=2) print(" batch=1 : %.4f s" % t1) print(" batch=2 : %.4f s" % t2) print(" 时间比 : %.4f" % (t2 / t1)) print(" → CPU 上 BLAS 的 gemm 是算术受限的,所以比值贴近 2;") print(" GPU 上 batch=1 时 SM 常常没填满,真实墙钟比会明显小于 2,") print(" 理论样本算术量约翻倍;吞吐与峰值显存应在实际硬件测量。") print("\n[L3] 换算成端到端") for steps in (20, 30, 50): print(" %2d 步 DDIM:CFG off 共 %2d 次前向 | CFG on 共 %2d 次前向" % (steps, steps, 2 * steps)) print(" 蒸馏类模型(LCM / SDXL-Turbo 那一支)把 w 变成网络的一个输入,") print(" 只跑一遍前向 —— 这就是为什么「干掉 CFG」是提速的第一优先级。") print("\n" + "=" * 78) if __name__ == "__main__": main() make_figures.py # -*- coding: utf-8 -*- """画本文的四张图。数据源全部来自 cfg_gmm.py / cfg_lab.py 的真实输出,不另造数。 cfg_geometry.png 引导把样本推到了哪:真条件 / 倾斜目标 / CFG 实际生成 w_sweep.png 纯度涨了多少、代价涨了多少(双轴) interval.png 引导只在某个噪声区间开,性价比差多少 cost_ledger.png 那个 2x 的账本,以及 w 对误差的放大 运行: /usr/local/bin/python3 make_figures.py 依赖: numpy, matplotlib(字体 PingFang SC) 注意:matplotlib 的 mathtext 标签一律用 raw 字符串,且反斜杠后面只能跟字母 ——源码会被 sync-code 原样搬进文章附录,反斜杠后面跟非字母会被体检器判成转义污染。 """ import os import numpy as np import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import cfg_gmm as G import cfg_lab as L 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) C_MAIN, C_ALT, C_GREEN, C_GRAY, C_PURPLE = ("#2563eb", "#dc2626", "#059669", "#6b7280", "#7c3aed") STEPS = 200 N = 20000 def smooth(H, sigma=1.6): """可分离高斯平滑(不依赖 scipy)。""" r = int(3 * sigma) xs = np.arange(-r, r + 1) k = np.exp(-0.5 * (xs / sigma) ** 2) k /= k.sum() H = np.apply_along_axis(lambda m: np.convolve(m, k, mode="same"), 0, H) H = np.apply_along_axis(lambda m: np.convolve(m, k, mode="same"), 1, H) return H def kde_on_grid(X, xs, sigma=2.0): """样本 -> 与 grid 同坐标的平滑密度(积分归一)。""" n = len(xs) lo, hi = xs[0], xs[-1] idx = np.clip(((X[:, 0] - lo) / (hi - lo) * n).astype(int), 0, n - 1) idy = np.clip(((X[:, 1] - lo) / (hi - lo) * n).astype(int), 0, n - 1) H = np.zeros((n, n)) np.add.at(H, (idy, idx), 1.0) H = smooth(H, sigma) return H / H.sum() def fig_geometry(): """三张等高线:真实条件 / 倾斜目标 / CFG 实际生成。""" rng, Xref = G.base_rng() # 与 cfg_gmm.main 同一条随机流 xs = np.linspace(-4.2, 4.2, 200) GX, GY = np.meshgrid(xs, xs) P = np.column_stack([GX.ravel(), GY.ravel()]) def dens(fn): d = np.exp(fn(P)) d = d.reshape(len(xs), len(xs)) return d / d.sum() D_cond = dens(lambda X: G.logp_cond(X)) D_un = dens(lambda X: G.logp_data(X)) w = 7.5 lp_t = G.tilted_logpdf(P, w) lp_t = lp_t - lp_t.max() D_tilt = np.exp(lp_t).reshape(len(xs), len(xs)) D_tilt /= D_tilt.sum() Xg = G.ddim_sample(rng.standard_normal((N, 2)), steps=STEPS, w=w) D_cfg = kde_on_grid(Xg, xs, sigma=2.2) mean_ref = Xref.mean(0) mean_cfg = Xg.mean(0) lvl_un = np.geomspace(D_un.max() * 1e-4, D_un.max(), 6) fig, axes = plt.subplots(1, 3, figsize=(15.2, 5.0)) panels = [ ("真实条件分布 $p(x|c)$", D_cond, mean_ref, C_MAIN), ("倾斜目标 $p(x|c)^{w}p(x)^{1-w}$", D_tilt, None, C_PURPLE), ("CFG 实际采出的分布", D_cfg, mean_cfg, C_ALT), ] for ax, (title, D, mean, col) in zip(axes, panels): ax.contour(xs, xs, D_un, levels=lvl_un, colors=[C_GRAY], linewidths=0.7, alpha=0.55) lvl = np.geomspace(D.max() * 2e-3, D.max(), 8) ax.contourf(xs, xs, D, levels=lvl, cmap="Blues", alpha=0.85) ax.contour(xs, xs, D, levels=lvl, colors=[col], linewidths=1.1) ax.plot(G.MU[:, 0], G.MU[:, 1], "k+", ms=11, mew=1.8) ax.plot(mean_ref[0], mean_ref[1], "o", color=C_MAIN, ms=9, markeredgecolor="white", mew=1.6) if mean is not None and not np.allclose(mean, mean_ref): ax.plot(mean[0], mean[1], "D", color=C_ALT, ms=8, markeredgecolor="white", mew=1.6) ax.annotate("", xy=mean, xytext=mean_ref, arrowprops=dict(arrowstyle="->", color=C_ALT, lw=2.0)) ax.set_title(title, fontsize=12) ax.set_xlim(xs[0], xs[-1]) ax.set_ylim(xs[0], xs[-1]) ax.set_aspect("equal") ax.grid(alpha=0.15) ax.set_xlabel("$x_1$") ax.set_ylabel("$x_2$") axes[0].text(0.02, 0.03, "灰色细线 = 无条件数据密度\n蓝色圆点 = 真实条件均值", transform=axes[0].transAxes, fontsize=8.5, color=C_GRAY, va="bottom") d_bias = float(np.linalg.norm(mean_cfg - mean_ref)) axes[2].text(0.02, 0.03, "红色菱形 = 生成分布均值\n离真实条件均值 %.2f(数据每维 std 约 0.91)" % d_bias, transform=axes[2].transAxes, fontsize=8.5, color=C_ALT, va="bottom") fig.suptitle(r"引导强度 $w=7.5$:样本被推到了哪(同一坐标系,网格 $[-4.2,4.2]^2$)", fontsize=13) fig.tight_layout(rect=[0, 0, 1, 0.94]) fig.savefig(os.path.join(FIGDIR, "cfg_geometry.png"), dpi=130) plt.close(fig) print(" cfg_geometry.png 类心偏移 %.4f" % d_bias) def fig_w_sweep(): """纯度涨了多少、代价涨了多少。""" rng, Xref = G.base_rng() # 与 cfg_gmm.main 同一条随机流 ws = np.array([0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 5.0, 7.5, 10.0, 15.0, 20.0, 25.0]) x_T0 = rng.standard_normal((N, 2)) rows = [G.summarize(G.ddim_sample(x_T0, steps=STEPS, w=float(w)), Xref) for w in ws] pur = np.array([r["purity"] for r in rows]) div = np.array([r["div"] for r in rows]) typ = np.array([r["typicality"] for r in rows]) bias = np.array([r["bias"] for r in rows]) fig, ax = plt.subplots(figsize=(9.6, 5.6)) ax.plot(ws, pur, "o-", color=C_MAIN, lw=2.2, ms=5, label="软纯度 $E[p(c|x)]$") ax.plot(ws, div, "s-", color=C_GREEN, lw=2.2, ms=5, label="多样性比(生成方差/真条件)") ax.axhline(0.7388, color=C_GRAY, ls=":", lw=1.2) ax.text(0.6, 0.752, "真实条件样本 0.7388", color=C_GRAY, fontsize=9) ax.set_xlabel(r"引导强度 $w$") ax.set_ylabel("纯度 / 多样性比", color="black") ax.set_ylim(0.35, 1.35) ax.grid(alpha=0.25) ax.legend(loc="center left", fontsize=9.5) ax2 = ax.twinx() ax2.plot(ws, typ, "^--", color=C_ALT, lw=2.2, ms=5, label=r"典型度 $\Delta\log p_{\mathrm{data}}$(nats)") ax2.plot(ws, bias, "v--", color=C_PURPLE, lw=2.0, ms=5, label="类心偏移") ax2.set_ylabel("典型度 / 类心偏移(越负/越大越糟)") ax2.legend(loc="center right", fontsize=9.5) ax.axvline(1.0, color=C_GRAY, lw=1.0, alpha=0.6) ax.axvline(7.5, color=C_ALT, lw=1.2, ls="-.", alpha=0.8) ax.text(7.9, 1.30, "SD 系列默认 $w=7.5$", color=C_ALT, fontsize=9.5) ax.text(1.1, 1.28, "$w=1$ 就是纯条件模型", color=C_GRAY, fontsize=9) fig.suptitle("纯度每涨一点,样本就离真实数据远一点(DDIM %d 步,N=%d)" % (STEPS, N), fontsize=12.5) fig.tight_layout(rect=[0, 0, 1, 0.95]) fig.savefig(os.path.join(FIGDIR, "w_sweep.png"), dpi=130) plt.close(fig) print(" w_sweep.png w=7.5: 纯度 %.4f / 多样性 %.4f / 典型度 %.4f / 偏移 %.4f" % (rows[8]["purity"], rows[8]["div"], rows[8]["typicality"], rows[8]["bias"])) def fig_interval(): """帕累托前沿:花多少「离数据流形的距离」,买多少「条件纯度」。""" rng, Xref = G.base_rng() # 与 cfg_gmm.main 同一条随机流 x_T0 = rng.standard_normal((N, 2)) base = G.summarize(G.ddim_sample(x_T0, steps=STEPS, w=1.0), Xref) def run(lo, hi): m = G.summarize(G.ddim_sample(x_T0, steps=STEPS, w=7.5, interval=(lo, hi)), Xref) gain = m["purity"] - base["purity"] cost = -(m["typicality"] - base["typicality"]) return gain, cost, m fig, ax = plt.subplots(figsize=(9.8, 6.6)) # 等性价比参考线(双对数坐标下是斜率 1 的直线);标签放在可见范围内 xr = np.logspace(-2.4, 0.6, 60) for k in (0.1, 0.5, 1.0, 5.0): ax.plot(xr, k * xr, "--", color=C_GRAY, lw=0.8, alpha=0.65) x_lab = 0.30 / k if 6e-3 <= x_lab <= 3.0: ax.text(x_lab, k * x_lab * 1.22, "性价比 %.1f" % k, fontsize=8, color=C_GRAY, ha="center", va="bottom", clip_on=True) # A 族:区间 = [0, hi],即只在链的后半(低噪声)开 his = (0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.8, 1.0) pts_a = [run(0.0, hi) for hi in his] ax.plot([p[1] for p in pts_a], [p[0] for p in pts_a], "-o", color=C_GREEN, lw=2.2, ms=8, markeredgecolor="white", mew=1.4, label=r"只在低噪声段开:区间 $[0,h]$") for hi, (g, c, m) in zip(his, pts_a): if hi in (0.1, 0.2, 0.5, 1.0): ax.annotate("$h$=%.1f\n性价比 %.2f" % (hi, g / c), (c, g), textcoords="offset points", xytext=(-64, 6), fontsize=9, color=C_GREEN) print(" 区间 [0.0,%.1f] 纯度增益 %+.4f 典型度代价 %.4f 性价比 %.3f" % (hi, g, c, g / c)) # B 族:区间 = [lo, 1],即只在链的前半(高噪声)开 los = (0.9, 0.8, 0.7, 0.6, 0.5, 0.4, 0.3, 0.2, 0.1, 0.0) pts_b = [run(lo, 1.0) for lo in los] ax.plot([p[1] for p in pts_b], [p[0] for p in pts_b], "-s", color=C_MAIN, lw=2.2, ms=7, markeredgecolor="white", mew=1.4, label=r"只在高噪声段开:区间 $[l,1]$") for lo, (g, c, m) in zip(los, pts_b): if lo in (0.9, 0.5, 0.0): ax.annotate("$l$=%.1f\n性价比 %.2f" % (lo, g / c), (c, g), textcoords="offset points", xytext=(10, -14), fontsize=9, color=C_MAIN) print(" 区间 [%.1f,1.0] 纯度增益 %+.4f 典型度代价 %.4f 性价比 %.3f" % (lo, g, c, g / c)) ax.set_xscale("log") ax.set_yscale("log") ax.set_xlabel(r"代价:典型度损失 $-\Delta\log p_{\mathrm{data}}$(nats)") ax.set_ylabel(r"收益:软纯度增益") ax.set_xlim(6e-3, 4.0) ax.set_ylim(6e-3, 0.4) ax.grid(alpha=0.25, which="both") ax.legend(loc="lower right", fontsize=10) fig.suptitle(r"引导区间的帕累托前沿($w=7.5$,DDIM %d 步)" % STEPS, fontsize=12.5) fig.tight_layout(rect=[0, 0, 1, 0.955]) fig.savefig(os.path.join(FIGDIR, "interval.png"), dpi=130) plt.close(fig) print(" interval.png") def fig_cost_ledger(): """2x 账本 + w 对误差的放大。""" rng, _ = G.base_rng() # 与 cfg_gmm.main 同一条随机流 mac0, act0, _ = L.ledger(with_cfg=False) mac1, act1, _ = L.ledger(with_cfg=True) fig, axes = plt.subplots(1, 3, figsize=(15.0, 4.6)) # (a) 账本 ax = axes[0] labels = ["单步 MAC", "层间张量\n(fp16)", "50 步的前向\n次数"] v0 = [mac0 / 1e11, act0 / 1024 ** 3, 50.0] v1 = [mac1 / 1e11, act1 / 1024 ** 3, 100.0] x = np.arange(3) ax.bar(x - 0.19, v0, 0.36, color=C_MAIN, label="CFG off") ax.bar(x + 0.19, v1, 0.36, color=C_ALT, label=r"CFG on($w>1$)") for i, (a, b) in enumerate(zip(v0, v1)): ax.text(i - 0.19, a, "%.2f" % a, ha="center", va="bottom", fontsize=8.5) ax.text(i + 0.19, b, "%.2f" % b, ha="center", va="bottom", fontsize=8.5) ax.set_xticks(x) ax.set_xticklabels(labels, fontsize=9.5) ax.set_yscale("log") ax.set_ylabel("(MAC 单位 1e11,显存单位 GB)") ax.set_title("那个 2 倍:定义级精确", fontsize=12) ax.legend(fontsize=9) ax.grid(alpha=0.2, axis="y") # (b) 误差放大 ax = axes[1] ws = np.array([1.0, 2.0, 3.0, 5.0, 7.5, 10.0, 15.0, 20.0]) probe = G.sample_data(rng, 20000) for t, col in ((1000, C_MAIN), (200, C_GREEN), (10, C_ALT)): a = G.abar(t) Xt = np.sqrt(a) * probe + np.sqrt(1 - a) * rng.standard_normal(probe.shape) ec = G.weak_eps_hat(Xt, t, G.LOGW_COND) - G.eps_hat(Xt, t, G.LOGW_COND) eu = G.weak_eps_hat(Xt, t, G.LOGW_UNCOND) - G.eps_hat(Xt, t, G.LOGW_UNCOND) base = np.linalg.norm(ec, axis=1).mean() rat = [np.linalg.norm(w * ec + (1 - w) * eu, axis=1).mean() / base for w in ws] ax.plot(ws, rat, "o-", color=col, lw=2.0, ms=5, label="实测 $t=%d$" % t) ax.plot(ws, 2 * ws - 1, "k--", lw=1.6, label=r"最坏上界 $2w-1$") ax.set_xlabel(r"引导强度 $w$") ax.set_ylabel("去噪误差被放大的倍数") ax.set_title("两个分支的误差也一起被放大", fontsize=12) ax.legend(fontsize=9) ax.grid(alpha=0.25) # (c) 幅度膨胀 ax = axes[2] Xb = G.sample_data(rng, 8000, cls=G.TARGET_CLS) a = G.abar(300) Xp = np.sqrt(a) * Xb + np.sqrt(1 - a) * rng.standard_normal(Xb.shape) e_c = G.eps_hat(Xp, 300, G.LOGW_COND) e_u = G.eps_hat(Xp, 300, G.LOGW_UNCOND) wgrid = np.linspace(0, 25, 120) infl = np.array([(e_u + w * (e_c - e_u)).std() / e_c.std() for w in wgrid]) ax.plot(wgrid, infl, "-", color=C_PURPLE, lw=2.4) for w in (1.0, 7.5, 15.0): v = float((e_u + w * (e_c - e_u)).std() / e_c.std()) ax.plot([w], [v], "o", color=C_ALT, ms=7) ax.annotate("$w$=%.1f:%.2f 倍" % (w, v), (w, v), textcoords="offset points", xytext=(8, -12), fontsize=9) ax.axhline(1.0, color=C_GRAY, ls=":", lw=1.2) ax.set_xlabel(r"引导强度 $w$") ax.set_ylabel(r"$\mathrm{std}(\varepsilon_{\mathrm{guided}})/\mathrm{std}(\varepsilon_{\mathrm{cond}})$") ax.set_title("预测幅度膨胀(过曝的机制)", fontsize=12) ax.grid(alpha=0.25) fig.tight_layout() fig.savefig(os.path.join(FIGDIR, "cost_ledger.png"), dpi=130) plt.close(fig) print(" cost_ledger.png MAC %.3e -> %.3e" % (mac0, mac1)) def main(): print("画图(数据源:cfg_gmm.py / cfg_lab.py 的真实输出)") fig_geometry() fig_w_sweep() fig_interval() fig_cost_ledger() print("输出目录:%s" % FIGDIR) if __name__ == "__main__": main()
2026年09月28日
2 阅读
0 评论
0 点赞
2026-09-27
AIGC 基本功|从 DDIM 到高阶采样器-DDIM
从 DDIM 到高阶采样器 所属方向:生成范式 | 难度:进阶 | 前置知识:DDPM 的训练目标与采样流程(ddpm)、扩散过程的前向与反向推导(diffusion_math) 关键词:DDIM、DPM-Solver、确定性采样、二阶采样、步数、NFE 01. 为什么需要它 先给两个数字,都来自文末附录里能直接跑的脚本。 数字一:同样的精度,评估次数差 12.8 倍。 在一个二维八分量高斯混合上跑(加噪之后仍是高斯混合,所以 score 有闭式解,见 04 节),把生成终点到高精度参考解的 RMS 距离压到 1e-2: 采样器 压到 1e-2 需要的模型评估次数 压到 1e-3 DDIM(时间步等间隔) 256 > 384 DPM-Solver-2(单步二阶) 63 191 DPM-Solver++ 2M(多步二阶) 40 128 DPM-Solver++ 3M(多步三阶) 20 64 DDIM 要 256 次评估,3M 只要 20 次。这不是"实现优化"级别的差距,是换了一套数学带来的差距——同一个模型、同一份起点、同一份随机数,只是步法不同。生产环境里 20 步出图还是 250 步出图,是"能上"和"不能上"的区别。 数字二:换个采样器,缓存策略就得重调。 这一条更隐蔽,也更容易踩。所有"复用上一步特征"的加速(DeepCache、各种 step-cache、特征缓存)都在做一个判断:这一步的模型输入和上一步差不多,就直接搬上一步的中间结果。那"差不多"到底是多少?在同一个模型、同一条 50 步轨迹上,量最后十步"模型输入相对变化": 采样器 最后 10 步输入相对变化 DDIM(η=0) 0.0197 η=0.5 0.0667 DDPM 祖采样(η=1) 0.1294 η=1 的抖动是 η=0 的 6.6 倍。也就是说,你在 DDIM 上把阈值调成"输入变化小于 2% 就复用",切到 DDPM 祖采样后这条曲线整段都在 13% 附近,缓存一次都不会命中,加速方案静默失效;反过来,为了吃掉 η=1 的抖动把阈值放宽到 15%,切回 DDIM 后就会在真正需要重算的步上复用,画面细节直接掉。这不是调参没调好,这是两套采样器的轨迹性质不同。 所以这篇文章要回答的是一个问题:采样器到底是什么? 答案不是"一种跳步策略",而是一条常微分方程的离散格式。一旦这么看,两件事立刻变得可算:步法是多少阶、达到给定精度要花多少次模型评估(NFE,number of function evaluations)。而 DDIM——这个被无数人当成"DDPM 的加速版"的东西——恰好是这条 ODE 上最低的一阶格式。下面会把它验到 1e-14。 02. 最小可用理解 三句话讲完: DDIM 不是"步幅更大的 DDPM"。 DDPM 的采样每步都要掷一份新噪声,DDIM 的贡献是证明:同一族前向过程可以用非马尔可夫的方式重新构造,只要边缘分布 $q(x_t|x_0)$ 不变,反向过程就有自由度——其中一个极端是完全不掷噪声。所以差别不在步幅,在随机项。 不掷噪声之后,采样变成解 ODE。 这个 ODE 叫概率流 ODE,它的解是确定性的:同一个起点永远给同一个终点。DDIM 的一步迭代,就是这条 ODE 的一个显式一阶格式(用当前点的信息走一整步)。用二维真轨迹看最直观: 这张图要看的是:蓝色的 η=0 轨迹是一条平滑曲线,橙色 η=1 轨迹在同一份随机数下走成了折线——每一步都被注入的方向改变。"同一个采样器家族"的两种极端,轨迹的几何性质完全不同,缓存类加速盯的正是这个几何性质。 既然它是 ODE,就有阶数,就有"性价比最高的一次评估花在哪"的问题。 把状态量从 $t$ 换成 $\lambda=\log(\alpha_t/\sigma_t)$(信噪比的对数)之后,一阶格式的误差正比于步长、二阶格式正比于步长的平方。于是同样的精度,高两阶的方法能省一个数量级的评估——这就是 DPM-Solver 系列的全部立足点。 03. 数学推导 3.1 把 DDIM 的一步拆开 DDPM 的采样式(祖先采样)长这样:先估干净图,再缩放,再加一份后验噪声。 $$x_{t-1} = \frac{1}{\sqrt{\alpha_t}}\Big(x_t - \frac{\beta_t}{\sqrt{1-\bar\alpha_t}}\epsilon_\theta(x_t,t)\Big) + \sqrt{\tilde\beta_t}z$$ DDIM 论文式 (12) 换了一个组织方式。令 $\hat x_0$ 是这一时刻对干净图的估计、$\epsilon_\theta$ 是模型输出的噪声预测,两者由 $\hat x_0 = (x_t-\sqrt{1-\bar\alpha_t}\epsilon_\theta)/\sqrt{\bar\alpha_t}$ 相互确定,则一步写成三项: $$x_{t-1} = \underbrace{\sqrt{\bar\alpha_{t-1}}\hat x_0}_{\text{回数据的方向}} + \underbrace{\sqrt{1-\bar\alpha_{t-1}-\tilde\sigma_t^2}\epsilon_\theta}_{\text{沿当前噪声方向}} + \underbrace{\tilde\sigma_t z}_{\text{新掷的噪声}}$$ 这里第二项必须使用目标时刻的总噪声方差 $1-\bar\alpha_{t-1}$,新注入噪声为 $$\tilde\sigma_t = \eta\sqrt{\frac{1-\bar\alpha_{t-1}}{1-\bar\alpha_t}}\sqrt{1-\frac{\bar\alpha_t}{\bar\alpha_{t-1}}}$$ 每一项在干什么。 第一项是"这一时刻认为的干净图"乘上新时刻的 $\sqrt{\bar\alpha}$,也就是把估计值放到目标噪声水平上。第二项补偿两者的差额:前向过程是 $x_t=\sqrt{\bar\alpha_t}x_0+\sqrt{1-\bar\alpha_t}\epsilon$,要从 $x_t$ 走到 $x_{t-1}$,得把噪声幅度从 $\sqrt{1-\bar\alpha_t}$ 缩到 $\sqrt{1-\bar\alpha_{t-1}}$,这个缩放只能沿着"当前噪声方向"做——$\sqrt{1-\bar\alpha_{t-1}-\tilde\sigma_t^2}$ 正好是这两者的勾股差。第三项才是唯一带随机性的部分。 $\eta$ 是那个旋钮。 $\eta=0$ 时第三项消失,$\tilde\sigma_t=0$,迭代完全确定;$\eta=1$ 时 $\tilde\sigma_t$ 恰好等于 DDPM 的后验标准差,第二、三项合起来正好还原 DDPM 的加噪——在相同的完整时间表、fixed-small 方差、预测器及随机数下,$\eta=1$ 与该 DDPM 更新代数等价;跳步时应比较对应重排时间表的祖先采样,而不是原始 1000 步链。04 节会把这两条都量到浮点误差量级。 为什么可以有这个旋钮。 因为 DDIM 用的前向过程不是马尔可夫链:用 $q_\eta(x_{t-1}\mid x_t,x_0)$ 与终点分布重新构造非马尔可夫联合过程,但所有边缘 $q(x_t|x_0)$ 一个都没变。训练目标只依赖边缘(噪声预测的回归目标就是 $\epsilon$),所以模型不用重训;而反向的可选空间变大了,$\eta$ 就是在里面挑一条路。 3.2 换成 ODE 视角 $\eta=0$ 时迭代没有随机项,它可以被看成一条 ODE 的离散格式。连续时间下,前向过程是 $dx = -\frac{1}{2}\beta(t)x\mathrm{d}t + \sqrt{\beta(t)}\mathrm{d}w$(第一项把 $x$ 往 0 拉,第二项持续注入白噪声),对应的概率流 ODE(与 SDE 共享所有边缘分布的那个确定性方程)是 $$\frac{dx}{dt} = -\frac{1}{2}\beta(t)\Big(x + \nabla_x\log p_t(x)\Big)$$ 两式一比就能看出这个方程的来历:把 SDE 的噪声项换成"噪声的均值流",也就是 score 项,随机性就被抽掉了,但每一时刻的边缘分布 $p_t(x)$ 一模一样。这就是为什么用同一条训练好的网络既能跑随机采样(SDE)、也能跑确定性采样(ODE)。 score 与噪声预测的关系是 $\nabla_x\log p_t(x) \approx -\epsilon_\theta(x,t)/\sqrt{1-\bar\alpha_t}$,代进去就得到一个只用模型输出的 ODE。关键一步是换坐标。 令 $\lambda = \log(\alpha_t/\sigma_t) = \frac{1}{2}\log\frac{\bar\alpha_t}{1-\bar\alpha_t}$(半 log-SNR)。 这个坐标为什么叫"半":信噪比本身就是 $\alpha_t^2/\sigma_t^2 = \bar\alpha_t/(1-\bar\alpha_t)$,取对数再除以二,得到的是"振幅比"的对数——也就是 $x_0$ 与 $\epsilon$ 两个分量在 $x_t$ 里的相对尺度的对数。$\lambda$ 越大代表越干净:$t=1000$ 时 $\bar\alpha_t\approx 4\times10^{-5}$,$\lambda\approx-5$;$t=1$ 时 $\bar\alpha_t\approx0.9999$,$\lambda\approx+4.6$。整条轨迹在 $\lambda$ 上只跨了不到 10 个单位,而 $t$ 跨了 1000 个单位——这就是"难度不是按 $t$ 均匀分布"的定量说法。 为什么要换?因为在这条 ODE 里,步长的含义由 $\lambda$ 决定:$t$ 上均匀的一步,落在 $\lambda$ 上的长度可能差几十倍。$\lambda$ 坐标把 ODE 拉成接近常数系数的形式($\alpha_\lambda$、$\sigma_\lambda$ 随 $\lambda$ 的变化是光滑的 sigmoid 型),指数积分才有干净的闭式。换成 $\lambda$ 之后,一步从 $\lambda_s$ 走到 $\lambda_t$(约定 $h=\lambda_t-\lambda_s>0$,因为我们从噪声走到干净)的精确解是 $$x_t = \frac{\sigma_t}{\sigma_s}x_s + \alpha_t\int_{0}^{h} e^{-(h-u)}\hat x_0(\lambda_s+u)du$$ 这个式子值得停一下:它说明整步的解由"$\hat x_0$ 沿着 $\lambda$ 的变化曲线"加权积分决定,权重是 $e^{-(h-u)}$——越靠近区间右端(越干净的那端)权重越大。验证一下:如果 $\hat x_0$ 是常数,积分给出 $\alpha_t(1-e^{-h})\hat x_0$,加上第一项正好是 $\alpha_t \hat x_0+\sigma_t \epsilon$,也就是"一步跳到位"的精确解。所以步法阶数的本质,就是用多少个点上的 $\hat x_0$ 去逼近这条曲线。 3.3 阶数从哪来:泰勒展开 把 $\hat x_0(\lambda_s+u)$ 在 $u=0$ 处展开,记 $\hat x_0,\hat x_0',\hat x_0''$ 为该点的一至二阶导,逐项积出系数: $$x_t = \frac{\sigma_t}{\sigma_s}x_s + \alpha_t\Big[J_0\hat x_0 + J_1\hat x_0' + J_2\hat x_0'' + \cdots\Big]$$ $$J_0 = 1-e^{-h},\qquad J_1 = e^{-h}-1+h,\qquad J_2 = \frac{h^2-2h+2-2e^{-h}}{2}$$ 小 $h$ 展开看数量级:$J_0\approx h$、$J_1\approx h^2/2$、$J_2\approx h^3/6$。所以导数项每高一阶,整步误差就多一个 $h$。 一阶格式:只留 $J_0$,把 $\hat x_0$ 当成整步不变。局部误差 $O(h^2)$,全局 $O(h)$。 二阶格式:再加 $J_1\hat x_0'$,系数 $\varphi_2 = J_1/h$,用一次评估估出 $\hat x_0'$ 就行。 三阶格式:再加 $J_2\hat x_0''$,系数 $\varphi_3 = J_2/h^2$。 多步法(multistep)就是"不额外花评估,用前几步的 $\hat x_0$ 做差商估导数":设这一步的 $\lambda$ 为 $\lambda_{s_0}$,前两步在 $\lambda_{s_1}$、$\lambda_{s_2}$,令 $h_i=\lambda_{s_i}-\lambda_{s_{i+1}}$,则 $$D_{1,0} = \frac{h}{h_0}(\hat x_0^{s_0}-\hat x_0^{s_1}) \approx h\hat x_0',\qquad D_1 = D_{1,0} + \frac{r_0}{r_0+r_1}(D_{1,0}-D_{1,1})\approx h\hat x_0'$$ 上一步的 $O(h)$ 偏差被组合系数消掉了(这就是"两个差商加权成一阶导"的标准手法),所以 $D_1$ 对 $\hat x_0'$ 是二阶准确的。同理二阶导用二阶差商: $$D_2 = \frac{D_{1,0}-D_{1,1}}{r_0+r_1}$$ 这里有个必须预先说清的细节。 三个点上的二阶差商等于 $\hat x_0''/2$(这是差商定义直接给出的),所以 $D_2\approx h^2\hat x_0''/2$;而上面泰勒系数 $\varphi_3$ 是按 $h^2\hat x_0''$ 的量纲推的。两者一比差一个 2。这个 2 在 05 节会展开:上游实现里单步版带了它、多步版没带,我先用受控实验把"该不该带"量出来,再决定怎么在文章里说。 3.4 DDIM 就是一阶格式 把 03.1 的 DDIM 迭代($\eta=0$)和"只留 $J_0$"的一阶格式对照:$\eta=0$ 时第二项系数为 $\sqrt{1-\bar\alpha_{t-1}}$,第三项为 0,代入 $\hat x_0$ 的定义整理,得到 $$x_t = \frac{\sigma_t}{\sigma_s}x_s + \alpha_t(1-e^{-h})\hat x_0(\lambda_s)$$ 与一阶格式逐项相同。所以"DDIM 是一阶方法"不是比喻,是恒等式——04 节用两套独立实现的代码把它验到 1e-14 以内。 04. 代码实现 全部代码在文末附录,五个脚本、只依赖 numpy:oracle_gmm.py(实验台与尺子)、ddim_family.py(η 家族)、dpm_solver_lab.py(λ 坐标与高阶格式)、spacing_lab.py(步数摆法)、make_figures.py(配图)。下面的数字是它们的真实输出。 4.1 先造一把可靠的尺子 要在二维上量"采样器差多少",需要一个模型误差为零的环境,否则量到的是模型不行,不是步法不行。用八个高斯分量摆在一圈上,加噪之后仍是高斯混合,score 有闭式解——模型换成这个解析 score,误差就只剩离散化。oracle_gmm.py 自己先做两件事: A. Oracle 自检:解析 score vs 有限差分 t= 1 abar=0.999900 max|解析 - 差分| = 1.57e-08 t= 200 abar=0.659039 max|解析 - 差分| = 4.99e-11 t= 1000 abar=0.000040 max|解析 - 差分| = 2.14e-11 B. 端点检查:t=T 时协方差与 I 的最大差 = 3.834e-05 C. 尺子的噪声地板:SW1 真vs真 ×8 = 0.03055 ± 0.00452 D. 天花板:真样本当生成样本送进去 SW1 = 0.000000 第一项是"我写的解析 score 不是拍脑袋写的"(和有限差分对上);第二项确认 $t=T$ 时分布确实接近标准正态(否则起点选错了);第三项最重要:两批都是真样本时也能量出 0.03 的距离,这是抽样噪声的地板,后面低于该尺度的差别需要更多样本和置信区间确认,不能仅凭单次结果归因;第四项确认尺子无偏置。 4.2 最小实现:DDIM 的一步 ddim_family.py 的核心就是下面这个函数,变量名与 3.1 节一一对应: def ddim_step(x, abar_cur, abar_tgt, eps, z, eta): """从 abar_cur 走到 abar_tgt(噪声变小)。返回 (x_tgt, x0_hat)。""" a_c, s_c = alpha_sigma(abar_cur) # 当前时刻的 alpha_t / sigma_t a_t, s_t = alpha_sigma(abar_tgt) # 目标时刻的 alpha_t / sigma_t x0 = (x - s_c * eps) / a_c # 第一项用的 x0_hat var = (1 - abar_tgt) / (1 - abar_cur) * (1 - abar_cur / abar_tgt) sig_tilde = eta * np.sqrt(max(var, 0.0)) # 第三项:新掷噪声的幅度 direc = np.sqrt(max(s_t ** 2 - sig_tilde ** 2, 0.0)) # 第二项:勾股差 return a_t * x0 + direc * eps + sig_tilde * z, x0 最后一行就是 3.1 节那三项,一次加法写完,没有任何按时刻分支的特判。跑个小验证,看 $\tilde\sigma_t$ 在轨迹两端各占多少: _, ABAR = linear_schedule() ts, grid = uniform_t_grid(50) # 51 个 abar,grid[0] 最吵、grid[-1] = 1.0 x, zs = paired_noise(50, n=4, seed=4321) # 起点 + 每一步要用的随机数 print("x.shape =", x.shape, "| grid 长度 =", len(grid)) for i in (0, 25, 48): for eta in (0.0, 1.0): a_c, s_c = alpha_sigma(grid[i]) a_t, s_t = alpha_sigma(grid[i + 1]) var = (1 - grid[i+1]) / (1 - grid[i]) * (1 - grid[i] / grid[i+1]) sig_tilde = eta * np.sqrt(var) print(f"step {i:2d} eta={eta:.1f} sigma_tilde = {sig_tilde:.6f}" f" (该步总噪声幅度 {s_t:.4f})") 真实输出: x.shape = (4, 2) | grid 长度 = 51 step 0 eta=0.0 sigma_tilde = 0.000000 (该步总噪声幅度 1.0000) step 0 eta=1.0 sigma_tilde = 0.574284 (该步总噪声幅度 1.0000) step 25 eta=0.0 sigma_tilde = 0.000000 (该步总噪声幅度 0.9509) step 25 eta=1.0 sigma_tilde = 0.419845 (该步总噪声幅度 0.9509) step 48 eta=0.0 sigma_tilde = 0.000000 (该步总噪声幅度 0.0760) step 48 eta=1.0 sigma_tilde = 0.063819 (该步总噪声幅度 0.0760) 两个数值得盯一下。η=1 时 $\tilde\sigma_t$ 与该步总噪声幅度之比,在起点是 57%(0.574/1.000),到接近干净端的第 48 步反而是 84%(0.0638/0.0760)。 也就是说"加噪声"的代价在轨迹末端最重:那里 $\hat x_0$ 刚要把几个模式分辨开,一份占 84% 的新噪声就砸进去了。这是 4.6 节 η=0 处处最优的直接原因。 再看一步里三项谁大谁小(用 oracle score 走第 25 步): eta=1 第 25 步三项 RMS: 回数据 0.1742 | 沿噪声方向 0.8694 | 新掷噪声 0.7800 两个噪声项的模长比"回数据"项大四五倍。三项向量的范数反映更新的组成,不能用向量大小推出网络算力花在哪里,而不是往数据方向推——这也解释了为什么"阶数"值钱:阶数讲的正是怎么更准地把这一步搬完。 4.3 两次交叉验证 ddim_family.py 里故意写了两套互不相干的代码:一套照 DDPM Algorithm 2 的原式走 1000 步(一步不能跳),一套照 DDIM 式 (12) 走。喂同样的随机数: A. 终点逐元素最大差 = 3.109e-14 终点 RMS = 1.3520 3.1e-14 就是浮点累加误差的量级。η=1 的 DDIM 与 DDPM 祖采样是同一个算法,不是"近似"。 dpm_solver_lab.py 里同样对待 DDIM 与一阶指数积分器: A. max|DDIM(eta=0) - DPM-Solver-1| S= 20 uniform-t 8.549e-15 uniform-lambda 4.774e-15 S= 50 uniform-t 1.044e-14 uniform-lambda 1.088e-14 两套记号、两套时间网格,输出差在 1e-14。DDIM 是一阶指数积分器,这句到这里可以当结论用了。 4.4 量阶数:误差 vs NFE 有了参考解(多步三阶跑 4000 步,再用 2000 步自查,分辨率 1.16e-08)就能量阶数了。横轴必须是真实 NFE——脚本里用 Counter 包住模型,把调用次数数出来,而不是拿"步数 × 理论阶数"算。 这张图要看的是斜率:DDIM 的误差衰减阶约 1,二阶方法约 2,修正系数后的三阶方法接近 3;误差更小的曲线位于图的下方。纵向的差距就是"同样的评估次数,精度差多少个数量级"。 拟合出来: 采样器 拟合阶数 p 备注 DDIM(时间步等间隔) 0.97 一阶 DDIM(λ 等间隔) 0.99 一阶,换坐标不改阶数 DPM-Solver-2(单步) 1.90 二阶,每步 2 次评估 DPM-Solver++ 2M(多步) 1.98 二阶,每步 1 次评估 DPM-Solver++ 3M(原式) 2.07 应该是三阶,实测只有二阶 3M(三阶项系数 ×2 后) 2.91 接近三阶 3M 那两行是全篇唯一"和教科书不一致"的地方,见 05 节。 4.5 步数与质量的实测表 阶数是数值分析的语言,产品要的是"20 步够不够"。同一批配置换成 SW1(切片 Wasserstein-1,越小越好),换几份起点量出抖动: NFE DDIM(t) DDIM(λ) Solver-2 2M 3M(原式) 10 0.1292±0.006 0.2180±0.006 0.1910±0.007 0.0731±0.006 0.0680±0.006 20 0.0741±0.004 0.1122±0.004 0.0568±0.004 0.0317±0.006 0.0277±0.005 50 0.0420±0.007 0.0516±0.007 0.0335±0.006 0.0328±0.005 0.0327±0.005 噪声地板是 0.0301±0.0026。请只看 NFE=20 那一行:20 次评估时各采样器之间的差距(0.0277 到 0.1122)远大于抖动;到了 50 次评估,所有方法都贴着地板(0.033~0.052),这张表就再也排不出名次了——不是"大家都一样好",是当前样本量和重复次数难以稳定区分,增加样本与重复实验仍可提高分辨力。要排名次得回到 4.4 的确定性误差表。 4.6 η 到底该调多大 ddim_family.py 在 oracle 下扫一遍 η,再在欠拟合模型下扫一遍。所谓欠拟合是让模型以为数据是单个高斯(用真均值真协方差拟合),它在高噪声区几乎是对的、在低噪声区错得离谱——这是人为选择的一种误差结构,并不代表所有真实网络(也正好对应"训练时看到的是加噪数据分布"这件事)。 这张图要看的是两件事:一是曲线从左到右单调下降(步数越多越好,符合预期);二是本玩具大多数设置下 η=0 的 SW1 更低,不能推广为所有模型中确定性采样必胜,至少在有闭式解、且模型差得很有代表性的两种情况下都是错的。右图还给出一个诚实的例外:N=100 时 η=1 的 0.1830 略低于 η=0 的 0.1845,但差值 0.0015 远小于这批量的抖动,不该当成结论。 4.7 缓存为什么会被采样器带崩 这张图要看的是:三条线在干净端都会收敛(步长趋于 0),但在中间段整段差 6.6 倍。缓存阈值是按这条线定的,换采样器就是换这条线的量级。同一份数据里 $x_0$ 估计的位移更夸张:η=0 是 0.0371,η=1 是 0.2407(6.5 倍)。 4.8 评估预算 这张图要看的是柱子的相对高度,以及"步数"和"评估次数"不是一回事:DPM-Solver-2 每步要 2 次评估(先踩一步到中点、再走完),所以它的 NFE 是步数的两倍,这也是它虽然二阶却在低 NFE 段不占优的原因。 4.9 步数该摆在哪:收尾步不是免费的 阶数回答的是"步数变多时误差掉多快",没回答"步数摆在哪"。spacing_lab.py 固定 NFE=20、只改摆法,终点误差如下: 摆法 DDIM(一阶) 2M(二阶) 在 $t$ 上等间隔(本文的时间网格) 1.107e-01 5.285e-02 在 $\lambda$ 上等间隔 1.713e-01 2.897e-02 在 $\lambda$ 上等间隔、落回整数 $t$ 1.712e-01 2.894e-02 Karras / EDM($\rho=7$) 1.873e-01 7.038e-02 Karras($\rho=3$) 2.655e-01 3.397e-01 阶数和网格要配套:一阶方法在 $t$ 上等间隔最好,二阶方法在 $\lambda$ 上等间隔最好(差 1.8 倍)。而 EDM 那套 $\rho=7$ 的摆法在这里垫底——它不是不好,是它的 $\sigma$ 区间(EDM 原文是 80→0.002)和 VP 调度(这里只有 158→0.01,注意 EDM 的 $\sigma$ 是 $\sigma_t/\alpha_t$,不是 $\sqrt{1-\bar\alpha_t}$)不是一回事,照搬会把最吵的那一步拉得极长:NFE=20 时它最大的 $\lambda$ 间隔是 1.02,而 $\lambda$ 等间隔只有 0.51。 uniform-t 为什么赢,值得追问一层。 它把 $\lambda$ 的上界截在了 +1.76(NFE=20 时),剩下的干净端交给最后那一步"直接输出 $\hat x_0$"。这条捷径是不是免费的?不是。沿高精度参考轨迹量"在 $\lambda$ 处直接输出 $\hat x_0$ 会偏多少": $\lambda$ −0.23 +0.74 +1.71 +2.67 +3.64 提前收尾的偏差 5.24e-01 1.40e-01 4.40e-02 9.39e-03 1.63e-03 截得越早白送的偏差越大,而且这份偏差不随步数下降——它只取决于截在哪个 $\lambda$。于是 $\lambda_{max}$ 存在最优值(2M,NFE=20,只改截断点): $\lambda_{max}$ 1.0 2.0 2.5 3.0 4.61(走到底) 终点误差 9.82e-02 2.80e-02 1.71e-02 1.81e-02 2.90e-02 U 形,最优在 2.5,比"走到底"好 1.7 倍。这条也解释了 uniform-t 的行为为什么还不错:它的截断点随步数自动往右挪(NFE=20 时 1.76、50 时 2.55、100 时 3.1),等于一个廉价的自适应策略——但它是巧合,不是设计。 顺带一个反直觉的观察。 一步的局部误差带着一个 $\alpha_t$ 因子($\delta_i\approx\alpha_{t_i}\frac{h^2}{2}\hat x_0'$),所以误差权重是 $\alpha_t\,|d\hat x_0/d\lambda|$ 而不是单纯的导数。量下来:噪声端 1.02e-01、干净端 1.00e-01,几乎一样重,峰在 $\lambda\approx-0.2$(也就是 $\bar\alpha\approx0.4$ 的中间段)。常听到的"把步数往干净端挪"在本实验里不成立——难度是按 $\lambda$ 均匀铺开的,真正要避免的是"某一段特别长"(Karras 就是栽在这)。 05. 工业级实现对照 以 huggingface/diffusers 的 DPMSolverMultistepScheduler.step 为准(https://github.com/huggingface/diffusers/blob/main/src/diffusers/schedulers/scheduling_dpmsolver_multistep.py,以 2026-09 时的实现为准,上游会重构),最小实现与生产实现的差距集中在四处。 第一处:算法类型的开关是"预测什么"。 algorithm_type 有四个取值:dpmsolver 用噪声预测($\epsilon$-parameterization),dpmsolver++ 用数据预测($x_0$-parameterization);后缀 -sde-dpmsolver++ 则在同一条 ODE 上加回一个可控的噪声项。本实验包含噪声预测单步 Solver-2 与数据预测多步 Solver++。参数化会改变数值误差和引导稳定性,不能由训练 loss 的参数化权重图证明某条预测曲线总更光滑。 第二处:低步数时的降阶。 多步法开头没有历史可用,第一、二步必须降成一阶/二阶,lower_order_final 控制最后一步是否也降阶,final_sigmas_type="zero" 让最后一步直接落到 $\sigma=0$。这一条在我们的实现里对应"$\lambda_t=+\infty$ 时 $e^{-h}=0$、整式退化成 $x_t=\hat x_0$",不用特判。我专门验证过"开头降阶"是不是 3M 阶数不达标的原因:给第一步多加一次评估换成中点法(warm2s=True),实测阶数从 2.10 变成 2.10——一点没变,所以瓶颈不在起步。 第三处:时间步的排布。 timestep_spacing 有 leading、trailing、linspace(当前 DPMSolverMultistepScheduler 默认) 三种,同一个调度器换一种排布,20 步出图质量差很多。我们把它量化了:DDIM 在 $t$ 上等间隔时,NFE=191 的终点误差是 1.245e-02;换成 $\lambda$ 等间隔反而变差到 1.775e-02,达到 1e-2 分别需要 256 次和 384 次评估。这是"坐标选择比阶数更早起作用"的直接证据:一阶方法的误差由最大的那一段步长决定,均匀 $\lambda$ 的 $h$ 本来就是相等的;它改变了各时间段的评估密度,误差还由场的导数、传播与终端截断决定。所以别把"回到 λ 坐标"当成万能钥匙——它给高阶方法提供了干净的积分形式,但不自动给一阶方法更好的网格。 第四处,也是唯一需要标注存疑的:三阶项的系数。 我们逐行对齐 diffusers 的多步三阶更新后,实测阶数只有 2.07(4.4 节),不是 3。于是做了一个受控实验:不让 $\hat x_0$ 由 $x$ 决定,而是直接规定它是 $\lambda$ 的二次多项式,此时精确解可以用数值积分算出来,任何真正的三阶格式都应该一步算准。把区间长度 $h$ 从 0.4 缩到 0.025: 一阶: 4.17e-02 1.04e-02 2.57e-03 6.39e-04 1.59e-04 → h^2 ✓ 二阶: 8.16e-03 1.05e-03 1.32e-04 1.65e-05 2.07e-06 → h^3 ✓ 三阶-D2原式: 1.96e-03 2.46e-04 3.06e-05 3.81e-06 4.75e-07 → h^3 ✗ 三阶-D2乘2: 4.43e-13 6.23e-14 2.89e-15 1.11e-15 4.44e-16 → 一步精确 ✓ 一阶、二阶的收敛阶都对得上公式推导;三阶原式只有 $h^3$(等于二阶),把 $D_2$ 乘 2 之后直接掉到 1e-13~1e-16 的机器精度。我还用非均匀步长($h_0\neq h_1\neq h$)复验过一遍,结论不变:原式 8.3e-04、乘 2 之后 5.4e-14,所以这不是"等间隔时才碰巧"的巧合。 回头逐字核对上游源码(diffusers v0.30.0):单步版 scheduling_dpmsolver_singlestep.py 里写的是 D2 = 2.0 * (D1_1 - D1_0) / (r0 - r1),带了那个 2;多步版 scheduling_dpmsolver_multistep.py 里写的是 D2 = (1.0 / (r0 + r1)) * (D1_0 - D1_1),没带。两个版本共用同一个 $\varphi_3$ 系数 - (alpha_t * ((exp(-h) - 1.0 + h) / h**2 - 0.5)) * D2。 我把这件事的处理方式说清楚:这篇文章只报告我们自己的复现结果——逐行对齐上游的多步三阶实现、实测阶数 2.07;把 $D_2$ 乘 2 后升到 2.91,局部误差检验一步算准。至于这是上游实现的笔误、还是我对多步差商归一化的理解和原作者不同,我不下结论,已记入待决清单单独确认。对读者的实用含义是确定的:不要假设"用了 3M 就是三阶",阶数要自己量。 06. 代价与边界 省了什么。 DDIM 把每步的随机项去掉,换来三件事:同样的步数下误差更小(4.6 节的两个模型都验证了)、轨迹光滑(缓存加速可用)、以及固定初始噪声时可复现的输出。正则连续 ODE 流可逆,有限步 DDIM 或终步投影并不自动保证一一对应。 赔了什么。 路径随机性少了。 给定起点只有一条轨迹,但更换初始噪声仍产生多样本;确定性 ODE 不意味着输出分布没有多样性,也没有额外规定必须付出的批量成本。 高阶方法的稳定域有限。 显式格式的稳定性有上限,步长太大时高阶项不是"更准"而是"发散"。生产实现里低 NFE 时的降阶、以及当前多步调度器的 solver_type="midpoint" / "heun" 选项;它们不能和作者单步 API 中的 dpmsolver / taylor 选项混用都是在拿稳定性换名义阶数。 理论阶数要靠光滑性支撑。 我们的实验用的是解析 score,$\hat x_0(\lambda)$ 足够光滑,阶数才量得出来。真实网络输出的 $\hat x_0$ 带高频抖动,阶数通常要打折扣——这解释了为什么"3 阶"在实践中省下的没那么多。 什么时候不该用确定性采样器。 需要随机性作为正则的场景(例如低步数下用 SDE 采样换取更好的分布覆盖、或者需要"温度"调节多样性),$x_0$ 强约束会导致过平滑;以及任何依赖"步步重掷噪声"来做布朗桥类操作的训练/蒸馏流程。 什么时候值得上高阶。 评估预算在 20~50 次这个区间时最划算(4.8 节:NFE=20 时 3M 已经压到 1e-2,DDIM 需要 256 次);一旦预算到几百次,所有方法都进入渐进区,选最便宜的一阶反而更省心。 07. 经典论文脉络 DDIM,arXiv:2010.02502(2020):把 DDPM 的马尔可夫反向替换成非马尔可夫族,边缘不变、反向可选,给出 $\eta$ 旋钮。它的历史意义是把"采样"从随机过程问题变成确定性 ODE 问题。 DPM-Solver,arXiv:2206.00927(2022):把 ODE 换到 $\lambda$ 坐标做指数积分,给出二阶/三阶的单步格式与"约 10 步出图"的结果。核心贡献是 $\lambda$ 坐标与精确的指数积分系数,阶数第一次变得可算。 DPM-Solver++(2022):把参数化从 $\epsilon$ 换成 $x_0$,并补上多步变体(每步只要 1 次评估),这才是现在框架默认调度器的形状;同时给出引导采样的稳定化处理。 EDM,arXiv:2206.00364(2022):把"调度(噪声表)"和"采样器(ODE 积分格式)"彻底解耦,并给出把任何 $\sigma$ 上的网络包装成统一 ODE 的框架。看完 EDM 再回看 DDIM/DPM-Solver,会发现它们只是同一套 ODE 的三组离散格式。 一致性模型 / LCM(2023):把"多步 ODE 积分"压成"一步直接映射到 $x_0$",用蒸馏替代阶数。它的定位不是"更高阶的采样器",而是"不需要采样器"。 08. 常见误解 误解一:DDIM 就是步幅更大的 DDPM。 不是。DDIM 可以走 $t$ 上的任意子序列(这是它"能跳步"的前提,不是它的内容),它的内容是去掉随机项。反过来也成立:确定性 DDIM(η=0)也可以走 1000 步,仍不同于随机 DDPM;η=1 配合对应完整时间表时才有前文的等价关系。把这两件事分开之后,"为什么 DDIM 的 1000 步也不等于 DDPM"这个问题才有答案。 误解二:步数少了,就把 η 调大一点补回来。 实测反了。oracle 下 η=0 在 N=10/20/50/100 全部最优(N=10 时 0.1321 对 0.1888);欠拟合模型下同样单调(N=10 时 0.2679 对 0.3496)。直觉为什么错:加噪声确实让每步的 x̂₀ 估计被"抹匀"一点,但代价是它也抹掉了本该累积的方向信息;而噪声项还会让下一步的输入抖动变大(6.6 倍那个数字),把误差一路带下去。 误解三:阶数越高越省。 一是不一定真拿到那个阶(4.4 节 3M 实测 2.07);二是低 NFE 段高阶方法要"攒历史",前几步被迫降阶,实测 NFE=10 时 2M(0.0731)确实明显好于 DDIM(0.1292),但 3M(0.0680)比 2M 只好一点点,收益远小于"二阶到三阶"的名义差距。 误解四:SW1/FID 分不出来就没差别。 4.5 节里 50 次评估时全部方法都贴在 0.031 附近,看着"都一样"——那是因为尺子的噪声地板就是 0.0301。指标贴着地板时不代表方法等价,只代表这个指标失效了,此时要换确定性指标(有真值时用 RMS,我们的表就是这么打的)。 误解五:换到 λ 坐标只会更好。 不保证。λ 坐标便于推导指数积分,但选择 λ 均匀网格是另一件事;本文的一阶实验反而更偏好 t 网格:DDIM 在 $\lambda$ 等间隔网格上误差更大(NFE=191 时 1.775e-02,而 $t$ 等间隔是 1.245e-02)。网格与阶数要配套选。 09. 动手验证 五个脚本都能直接跑,几分钟内出结果。建议按这个顺序试: cd outputs/fundamentals_files/ddim_samplers/code /usr/local/bin/python3 oracle_gmm.py # ① 标定尺子(先量地板) /usr/local/bin/python3 ddim_family.py # ② η 旋钮与缓存抖动 /usr/local/bin/python3 dpm_solver_lab.py # ③ 阶数、预算、三阶项的受控检验 /usr/local/bin/python3 spacing_lab.py # ④ 步数摆法与收尾点 /usr/local/bin/python3 make_figures.py # ⑤ 画图 四个具体的改动实验,附我这里跑出来的结果: 把 η 从 0 改成 0.3:ddim_family.py 的 D 节会打印"输入相对变化",η=0 时最后 10 步是 0.0197,η=0.5 是 0.0667,η=1 是 0.1294。你会看到"加一点点噪声"的代价比想象中大。 把 dpm_solver_lab.py 里 dpm_pp_third 的 d2_scale 从 1.0 改成 2.0:C 节里 3M 的拟合阶数会从 2.07 变成 2.91,F 节的局部误差从 4.75e-07(最后一行)变成 4.44e-16。这一个常数就是"名义三阶"和"实测三阶"的分界。 把 DDIM 的网格从 uniform_t_grid 换成 uniform_lambda_grid:C 节里 DDIM 那两行的阶数不变(0.97 / 0.99),但误差整体上移,达到 1e-2 所需的 NFE 从 256 涨到 384。"同阶但更差"这件事,只有量出来才看得见。 在 spacing_lab.py 里把 grid_uniform_lambda_capped(20, cap) 的 cap 从 4.61 改成 2.5:E2 节那张 U 形表上,2M 的终点误差会从 2.90e-02 掉到 1.71e-02。这个数不用换模型、不用换采样器,只改一行网格。 10. 延伸阅读 ddpm:DDPM 的训练目标与采样流程。本篇的前置——$\hat x_0$ 与 $\epsilon_\theta$ 的互换关系、$\bar\alpha_t$ 的调度都在那里定义。 diffusion_math:扩散过程的前向与反向推导。概率流 ODE 与 SDE 的关系、参数化选择对误差权重的影响,是理解 05 节 algorithm_type 开关的基础。 vae_elbo:变分下界与重参数化。$\eta$ 旋钮的合法性来自"边缘分布不变",那套记号在那里建立。 cfg:分类器无关引导。引导会放大 ODE 右端的 Lipschitz 常数,采样步法要跟着调,是 06 节"稳定域"的直接后果。 flow_matching:流匹配与 Rectified Flow。把"直线路径"当作先验,是这条 ODE 主线的另一个分支。 附录:完整代码 09 节用到的脚本全文如下(oracle_gmm.py、ddim_family.py、dpm_solver_lab.py、spacing_lab.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 oracle_gmm.py # -*- coding: utf-8 -*- """采样器实验的地基:一个 score 有闭式解的二维高斯混合 + 一把「尺子」。 为什么必须先把这两件事钉死: 1. 比较采样器时,最大的干扰项是**模型误差**。同一个噪声网络,在 t=1 和 t=1000 上的误差可以差两个数量级,你根本分不清「这个采样器好」和「这个时间步好学」。 用闭式 score(oracle),模型误差恒为 0,剩下的全是离散化误差。 2. 比较采样器时,第二大的干扰项是**尺子本身的噪声**。切片 Wasserstein 距离 是用有限样本估的,两个「都对」的分布之间也能量出一个正数。这个正数就是 噪声地板,所有小于它的差别都是运气。 所以这个脚本干三件事:给出 oracle、给出尺子、量出地板。后面三个脚本全部 import 它,保证四处用到的目标分布、指标、随机数是同一套。 /usr/local/bin/python3 oracle_gmm.py 依赖:numpy, matplotlib(画图时才需要) """ import numpy as np T = 1000 # 训练时的扩散步数 D = 2 # 数据维度(二维才画得出轨迹,也才跑得起闭式 score) N_EVAL = 4000 # 评价用的样本数 N_DIRS = 512 # 切片 Wasserstein 的投影方向数 # ══════════════════════════════════════════════════════════════════ # 1. 噪声调度 # ══════════════════════════════════════════════════════════════════ def linear_schedule(T: int = T, b0: float = 1e-4, b1: float = 0.02): """DDPM / diffusers 默认的线性 beta 调度。 返回 abar 的长度是 T+1:abar[t] = prod_{s<=t} (1 - beta_s), abar[0] = 1 表示「完全干净」。这样下标和论文里的 t 直接对齐, 不用在 0-based / 1-based 之间来回换算。 """ betas = np.linspace(b0, b1, T) abar = np.concatenate([[1.0], np.cumprod(1.0 - betas)]) return betas, abar def alpha_sigma(abar_t: float): """给定 abar_t,返回 (alpha_t, sigma_t) = (sqrt(abar), sqrt(1-abar))。 这是全文唯一的「坐标」:x_t = alpha_t * x_0 + sigma_t * eps。 注意 alpha 是**信号尺度**,不是 DDPM 论文里的 alpha_t = 1 - beta_t。 """ return np.sqrt(abar_t), np.sqrt(1.0 - abar_t) # ══════════════════════════════════════════════════════════════════ # 2. 目标分布:8 个切向拉长的高斯,摆成一个环 # ══════════════════════════════════════════════════════════════════ def make_target(K: int = 8, radius: float = 1.8, tang: float = 0.30, rad: float = 0.05): """K 个分量均匀摆在一个半径 radius 的环上,每个分量沿切线方向拉长。 为什么不用各向同性:各向同性的高斯混合,加噪之后很快变成一个圆球, score 场几乎是线性的,一阶方法就能解到机器精度,分不出高下。 切向拉长之后,环上的 score 场有明显曲率,离散化误差才显出来。 """ ang = np.linspace(0.0, 2.0 * np.pi, K, endpoint=False) mu = np.stack([radius * np.cos(ang), radius * np.sin(ang)], axis=1) Sig = np.zeros((K, D, D)) for k, a in enumerate(ang): R = np.array([[np.cos(a), -np.sin(a)], [np.sin(a), np.cos(a)]]) Sig[k] = R @ np.diag([tang, rad]) @ R.T w = np.ones(K) / K return w, mu, Sig def sample_true(n: int, w, mu, Sig, rng): """从目标分布抽样。""" k = rng.choice(len(w), size=n, p=w) L = np.linalg.cholesky(Sig[k]) # (n, D, D) z = rng.standard_normal((n, D)) return mu[k] + np.einsum("nij,nj->ni", L, z) # ══════════════════════════════════════════════════════════════════ # 3. Oracle:加噪后分布的闭式 score,以及最优的 eps 预测 # ══════════════════════════════════════════════════════════════════ def _cov_at(abar_t: float, Sig): """q(x_t) 的第 k 个分量协方差:C_k = abar * Sig_k + (1 - abar) * I。""" return abar_t * Sig + (1.0 - abar_t) * np.eye(D) def logp_and_score(x, abar_t: float, w, mu, Sig): """返回 (log q_t(x), score_t(x)),x 形状 (n, D)。 q_t 是 K 个高斯的混合:第 k 个分量的均值是 sqrt(abar) * mu_k, 协方差是 abar * Sig_k + (1 - abar) * I。score 是这个混合分布的 对数梯度 = 各分量 score 的后验加权平均。 """ C = _cov_at(abar_t, Sig) # (K, D, D) inv = np.linalg.inv(C) ld = np.log(np.linalg.det(C)) r = x[:, None, :] - np.sqrt(abar_t) * mu[None] # (n, K, D) quad = np.einsum("nki,kij,nkj->nk", r, inv, r) lg = np.log(w)[None, :] - 0.5 * quad - 0.5 * ld[None, :] m = lg.max(axis=1, keepdims=True) p = np.exp(lg - m) Z = p.sum(axis=1, keepdims=True) logp = (m[:, 0] + np.log(Z[:, 0])) pi = p / Z # 后验分量权重 comp = -np.einsum("kij,nkj->nki", inv, r) # 每个分量的 score score = np.einsum("nk,nki->ni", pi, comp) return logp, score def eps_star(x, abar_t: float, w, mu, Sig): """MMSE 意义下最优的 eps 预测:eps* = -sigma_t * score_t(x)。 因为 x_t = alpha x_0 + sigma * eps,对 x_t 求 log 梯度会得到 grad log q_t = -E[eps | x_t] / sigma,所以最优的噪声预测就是 score 乘一个 -sigma。后面所有采样器都拿它当「模型」。 """ _, s = logp_and_score(x, abar_t, w, mu, Sig) return -np.sqrt(1.0 - abar_t) * s def x0_from_eps(x, abar_t: float, eps): """由 x_t 和 eps 反解 x_0:x_0 = (x_t - sigma_t * eps) / alpha_t。""" a, sg = alpha_sigma(abar_t) return (x - sg * eps) / a # ══════════════════════════════════════════════════════════════════ # 4. 尺子:切片 Wasserstein-1 + 平均对数密度 # ══════════════════════════════════════════════════════════════════ def make_dirs(n_dirs: int = N_DIRS, seed: int = 7): """预先抽好投影方向,全篇共用同一组。 这一条很重要:每次评价都重新抽方向的话,两组「同样好」的样本也会因为 方向不同量出不同的数,这部分方差完全是无谓的。方向固定之后, A 和 B 的差别就只来自样本本身。 """ th = np.random.default_rng(seed).standard_normal((n_dirs, D)) th /= np.linalg.norm(th, axis=1, keepdims=True) return th def sliced_w1(X, Y, TH=None, n_dirs: int = N_DIRS, rng=None): """切片 W1:在 n_dirs 个随机方向上算一维 W1,再取平均。 一维 W1 有闭式解:把两个样本集投影后排序,逐位相减取绝对值平均。 单位和数据同单位(这里是「平均要挪多远」),比 MMD^2 好解释。 """ if TH is None: TH = make_dirs(n_dirs) if rng is None else None if TH is None: th = rng.standard_normal((n_dirs, X.shape[1])) th /= np.linalg.norm(th, axis=1, keepdims=True) else: th = TH a = np.sort(X @ th.T, axis=0) b = np.sort(Y @ th.T, axis=0) return float(np.abs(a - b).mean()) def mean_logp(X, w, mu, Sig): """生成样本在真实分布下的平均对数密度。""" lp, _ = logp_and_score(X, 1.0, w, mu, Sig) return float(lp.mean()) def evaluate(X, Xref, w, mu, Sig, lp_ref, TH): """一次评价,返回 (SW1, dlogp)。 dlogp = 生成样本的平均 log 密度 − 真实样本的平均 log 密度。 0 表示完美;负数表示生成样本落在了真实分布的「空区」。 """ return sliced_w1(X, Xref, TH=TH), mean_logp(X, w, mu, Sig) - lp_ref # ══════════════════════════════════════════════════════════════════ # 5. 主流程:验 oracle、量尺子的噪声地板 # ══════════════════════════════════════════════════════════════════ def main(): rng = np.random.default_rng(20260927) betas, abar = linear_schedule() w, mu, Sig = make_target() print("=" * 68) print("A. Oracle 自检:解析 score vs 有限差分") print("=" * 68) for t in (1, 50, 200, 500, 1000): x = sample_true(8, w, mu, Sig, rng) _, s = logp_and_score(x, abar[t], w, mu, Sig) h = 1e-5 fd = np.zeros_like(s) for i in range(D): e = np.zeros(D) e[i] = h lp_p, _ = logp_and_score(x + e, abar[t], w, mu, Sig) lp_m, _ = logp_and_score(x - e, abar[t], w, mu, Sig) fd[:, i] = (lp_p - lp_m) / (2 * h) err = np.abs(fd - s).max() print(f" t={t:5d} abar={abar[t]:.6f} max|解析 - 差分| = {err:.2e}") print(" 差分误差在 1e-7 量级 → 解析 score 是对的(不是我拍脑袋写的)") # 顺便看一眼:oracle eps 在 t=1000 处还剩多少信息 print() print("=" * 68) print("B. 端点检查:t=T 时分布离标准正态有多远") print("=" * 68) for t in (200, 500, 800, 1000): C = _cov_at(abar[t], Sig) dev = np.abs(C - np.eye(D)).max() print(f" t={t:5d} abar={abar[t]:.3e} " f"各分量协方差与 I 的最大差 = {dev:.3e}") print(" 越接近 t=T,加噪后的混合越接近一个标准正态球") # ── 尺子的噪声地板 ── print() print("=" * 68) print("C. 尺子的噪声地板:两批「都对」的样本之间也能量出距离") print("=" * 68) Xref = sample_true(N_EVAL, w, mu, Sig, rng) lp_ref = mean_logp(Xref, w, mu, Sig) print(f" 参考样本: {Xref.shape}, 真实样本自己的平均 logp = {lp_ref:.4f}") TH = make_dirs() sw_floor, lp_floor = [], [] for r in range(8): Y = sample_true(N_EVAL, w, mu, Sig, np.random.default_rng(1000 + r)) sw_floor.append(sliced_w1(Y, Xref, TH=TH)) lp_floor.append(mean_logp(Y, w, mu, Sig)) print(f" SW1 真vs真 ×8: 均值 {np.mean(sw_floor):.5f} " f"标准差 {np.std(sw_floor):.5f} 最大 {np.max(sw_floor):.5f}") print(f" dlogp 真vs真 ×8: 均值 {np.mean(lp_floor) - lp_ref:+.4f} " f"标准差 {np.std(lp_floor):.4f}") print(" → 后面所有表格里,小于这个数的差别都是抽样噪声,不是采样器的功劳") # ── 一个有参照的好答案长什么样 ── print() print("=" * 68) print("D. 天花板:把真样本直接当生成样本送进去") print("=" * 68) print(f" SW1 = {sliced_w1(Xref, Xref, TH=TH):.6f}" f" dlogp = {mean_logp(Xref, w, mu, Sig) - lp_ref:+.4f}") print(" (恒等于 0,说明尺子本身没有偏置)") if __name__ == "__main__": main() ddim_family.py # -*- coding: utf-8 -*- """DDIM 家族:一个 eta 旋钮,把 DDPM 祖采样和 DDIM 串成一条线。 这个脚本回答三个问题: 1. **eta=1 的 DDIM 是不是就是 DDPM 祖采样?** 用两套独立的代码(一套写 DDPM Algorithm 2 的原式,一套写 DDIM 论文式 (12)),喂同样的随机数和 同样的起点,看输出能不能对到最后一位。能,那这条家族关系就不是传说。 2. **eta 到底该调多大?** 在 oracle score 下扫一遍:步数少的时候是不是 「加噪声更好」?在欠拟合模型下再扫一遍,看结论翻不翻。 3. **换采样器为什么缓存策略得重调?** 量每步「模型输入变了多少」: 确定性采样器的轨迹是光滑的,随机采样器每一步都往里砸一份新噪声。 /usr/local/bin/python3 ddim_family.py 依赖:numpy(+ 同目录的 oracle_gmm.py) """ import numpy as np from oracle_gmm import (T, linear_schedule, make_target, sample_true, eps_star, x0_from_eps, alpha_sigma, sliced_w1, mean_logp, make_dirs, N_EVAL, N_DIRS) # ══════════════════════════════════════════════════════════════════ # 0. 时间步网格与模型 # ══════════════════════════════════════════════════════════════════ def uniform_t_grid(S: int, T: int = T): """DDIM 论文 / diffusers 默认的「leading」取法:在 t 上等间隔。 返回 S+1 个 abar:起点 abar[T],终点 1.0(完全干净)。 """ ts = np.round(np.linspace(0.0, T, S + 1)).astype(int)[::-1] grid = np.array([1.0 if t == 0 else _ABAR[t] for t in ts], dtype=float) return ts, grid def paired_noise(S: int, n: int, seed: int = 1234): """给所有采样器**同一份**起点和同一批随机数。 这是整套实验最要紧的一条:如果每个配置各抽各的起点,配置之间就多了一份 抽样方差,量出来的差别里有多少是采样器的功劳根本说不清。配对之后, eta=0 和 eta=1 的终点之差只能来自算法本身。 """ r = np.random.default_rng(seed + S) x_init = r.standard_normal((n, 2)) zs = [r.standard_normal((n, 2)) for _ in range(S)] return x_init, zs def make_model(kind: str, w, mu, Sig): """两种「模型」,用来分离离散化误差和模型误差。 oracle : 用的就是真 score,模型误差 = 0,剩下的全是离散化误差。 gauss : 模型以为数据是一个单高斯(用真均值真协方差拟合), 也就是「容量不足以表达多峰」。高噪声区它几乎是对的, 低噪声区它错得离谱 —— 这正是真实网络的误差结构。 """ if kind == "oracle": return lambda x, abar: eps_star(x, abar, w, mu, Sig) if kind == "gauss": m = (w[:, None] * mu).sum(axis=0) S = ((w[:, None, None] * (Sig + np.einsum("ki,kj->kij", mu, mu))).sum(axis=0) - np.outer(m, m)) w1, mu1, S1 = np.array([1.0]), m[None, :], S[None, :, :] return lambda x, abar: eps_star(x, abar, w1, mu1, S1) raise ValueError(kind) # ══════════════════════════════════════════════════════════════════ # 1. DDIM 的一步(论文式 12,与 diffusers DDIMScheduler.step 逐项对齐) # ══════════════════════════════════════════════════════════════════ def ddim_step(x, abar_cur, abar_tgt, eps, z, eta): """从 abar_cur 走到 abar_tgt(abar_tgt > abar_cur,噪声变小)。 x_tgt = sqrt(abar_tgt) * x0_hat + sqrt(sigma_tgt^2 - sigma_tilde^2) * eps ← 指向 x_t 的方向项 + sigma_tilde * z ← 随机项 sigma_tilde = eta * sqrt( (1-abar_tgt)/(1-abar_cur) * (1 - abar_cur/abar_tgt) ) eta = 0 → 完全没有随机项,DDIM; eta = 1 → 随机项等于 DDPM 的后验标准差 beta_tilde。 """ a_c, s_c = alpha_sigma(abar_cur) a_t, s_t = alpha_sigma(abar_tgt) x0 = (x - s_c * eps) / a_c var = (1.0 - abar_tgt) / (1.0 - abar_cur) * (1.0 - abar_cur / abar_tgt) sig_tilde = eta * np.sqrt(max(var, 0.0)) direc = np.sqrt(max(s_t ** 2 - sig_tilde ** 2, 0.0)) return a_t * x0 + direc * eps + sig_tilde * z, x0 def ddpm_step(x, abar_cur, abar_tgt, eps, z, betas): """DDPM Algorithm 2 的原式,故意写成「另一套代码」用来交叉验证。 x_{t-1} = (x_t - beta_t / sqrt(1 - abar_t) * eps) / sqrt(alpha_t) + sqrt(beta_tilde_t) * z 只在相邻步(stride = 1)上成立,不能跳步。 """ abar_prev = abar_tgt beta_t = 1.0 - abar_cur / max(abar_prev, 1e-12) # = 1 - alpha_t beta_tilde = (1.0 - abar_prev) / (1.0 - abar_cur) * beta_t a_t = np.sqrt(max(abar_cur / max(abar_prev, 1e-12), 1e-12)) mean = (x - beta_t / np.sqrt(1.0 - abar_cur) * eps) / a_t return mean + np.sqrt(max(beta_tilde, 0.0)) * z def run_ddim(x, grid, model, eta, zs, track=False): """沿 grid 走一遍。grid[0] 最吵,grid[-1] = 1.0。 track=True 时额外返回每一步的 (x, x0_hat),D 节量「轨迹抖动」要用。 """ xs, x0s = [x.copy()], [] for i in range(len(grid) - 1): eps = model(x, grid[i]) z = zs[i] if zs is not None else 0.0 x, x0 = ddim_step(x, grid[i], grid[i + 1], eps, z, eta) if track: xs.append(x.copy()) x0s.append(x0.copy()) return x, (xs, x0s) if track else x0s # ══════════════════════════════════════════════════════════════════ # 2. 主流程 # ══════════════════════════════════════════════════════════════════ _BETAS, _ABAR = linear_schedule() def main(): rng = np.random.default_rng(20260927) w, mu, Sig = make_target() Xref = sample_true(N_EVAL, w, mu, Sig, rng) lp_ref = mean_logp(Xref, w, mu, Sig) ev_rng = np.random.default_rng(7) model = make_model("oracle", w, mu, Sig) # ── A. eta=1 的 DDIM 是不是 DDPM 祖采样 ── print("=" * 70) print("A. 交叉验证:DDIM(eta=1, stride=1) 是否等于 DDPM Algorithm 2") print("=" * 70) S = 1000 ts, grid = uniform_t_grid(S) n = 400 x_init = rng.standard_normal((n, 2)) zs = [rng.standard_normal((n, 2)) for _ in range(S)] x_ddpm = x_init.copy() for i in range(S): # 原式,一步步走完 1000 步 t = ts[i] x_ddpm = ddpm_step(x_ddpm, grid[i], grid[i + 1], model(x_ddpm, grid[i]), zs[i], _BETAS) x_ddim, _ = run_ddim(x_init.copy(), grid, model, eta=1.0, zs=zs) d = np.abs(x_ddpm - x_ddim).max() print(f" 两条独立实现,喂同样的 {S} 份随机数") print(f" 终点逐元素最大差 = {d:.3e} 终点 RMS = {np.sqrt((x_ddim**2).mean()):.4f}") print(f" → 差 {d:.1e},是浮点累加误差量级:eta=1 就是 DDPM 祖采样,不是「近似」") # ── B. eta 扫描(oracle score)── print() print("=" * 70) print("B. eta 扫描:步数越少,是不是越该加噪声?(oracle score,配对起点)") print("=" * 70) TH = make_dirs() floors = [sliced_w1(sample_true(N_EVAL, w, mu, Sig, np.random.default_rng(1000 + r)), Xref, TH=TH) for r in range(6)] print(f" 噪声地板 SW1 = {np.mean(floors):.4f} ± {np.std(floors):.4f}" f"(真实样本对真实样本),越接近它越好") NS = (10, 20, 50, 100) print(f" {'eta':>5} " + " ".join(f"{'N=' + str(s):>10}" for s in NS)) for eta in (0.0, 0.25, 0.5, 0.75, 1.0): row = [] for S in NS: ts, grid = uniform_t_grid(S) x_init, zs = paired_noise(S, N_EVAL) x, _ = run_ddim(x_init, grid, model, eta, zs) row.append(sliced_w1(x, Xref, TH=TH)) print(f" {eta:5.2f} " + " ".join(f"{v:10.4f}" for v in row)) # ── C. 换个欠拟合的模型,结论翻不翻 ── print() print("=" * 70) print("C. 同样的扫描,但把模型换成「以为数据是单高斯」") print("=" * 70) gm = make_model("gauss", w, mu, Sig) print(f" {'eta':>5} " + " ".join(f"{'N=' + str(s):>10}" for s in NS)) for eta in (0.0, 0.25, 0.5, 0.75, 1.0): row = [] for S in NS: ts, grid = uniform_t_grid(S) x_init, zs = paired_noise(S, N_EVAL) x, _ = run_ddim(x_init, grid, gm, eta, zs) row.append(sliced_w1(x, Xref, TH=TH)) print(f" {eta:5.2f} " + " ".join(f"{v:10.4f}" for v in row)) # ── D. 每步「模型输入变了多少」:缓存友好度 ── print() print("=" * 70) print("D. 轨迹抖动:每步「模型输入」变了百分之几(缓存复用看的就是它)") print("=" * 70) print(" 缓存类加速(按『这一步的输入和上一步差不多就复用上一步的特征』决策)") print(" 盯的是这个量。它一旦被噪声项垫住,原阈值就全废了。") print(f" {'eta':>5} {'最后10步 输入相对变化':>22} {'最后10步 x0_hat位移':>22}") S = 50 ts, grid = uniform_t_grid(S) x_init, zs = paired_noise(S, N_EVAL) out = {} for eta in (0.0, 0.5, 1.0): _, (xs, x0s) = run_ddim(x_init.copy(), grid, model, eta, zs, track=True) rel_x = [float(np.sqrt((((xs[i + 1] - xs[i]) ** 2).sum(1)).mean()) / np.sqrt(((xs[i] ** 2).sum(1)).mean())) for i in range(len(xs) - 1)] d_x0 = [float(np.sqrt((((x0s[i] - x0s[i - 1]) ** 2).sum(1)).mean())) for i in range(1, len(x0s))] out[eta] = (rel_x, d_x0) print(f" {eta:5.1f} {np.mean(rel_x[-10:]):22.4f} {np.mean(d_x0[-10:]):22.4f}") r0, r1 = out[0.0][0][-10:], out[1.0][0][-10:] print(f" → eta=1 的输入抖动是 eta=0 的 {np.mean(r1) / np.mean(r0):.1f} 倍;" f"按 eta=0 调出来的缓存阈值,换个采样器就完全不是一回事") if __name__ == "__main__": main() dpm_solver_lab.py # -*- coding: utf-8 -*- """从 DDIM 走到高阶:把时间轴换成 lambda,阶数就变成可以直接量出来的东西。 这个脚本做四件事: 1. **验证 DDIM 就是一阶 DPM-Solver。** 两套完全不同的代码(一套写 DDIM 论文式 12,一套写 lambda 空间的一阶指数积分器),同样的时间网格、 同样的模型,看输出是不是同一个数。是,那「DDIM 是一阶方法」就不是比喻。 2. **把阶数量出来。** 拿一条超高精度的参考解当真值,量各采样器终点到它的 距离。这个距离没有抽样噪声(起点配对、采样器确定),能干净地跨几个数量级, 于是「误差 ~ NFE^(-p)」里的 p 可以直接拟合出来。 3. **给出实用的步数-质量表。** 阶数是数值分析的语言,产品要的是 「20 步够不够」,所以再给一张 SW1 表(含噪声地板)。 4. **看高阶方法省下的评估次数到底落在哪。** /usr/local/bin/python3 dpm_solver_lab.py 依赖:numpy(+ 同目录的 oracle_gmm.py / ddim_family.py) """ import numpy as np from oracle_gmm import (T, linear_schedule, make_target, sample_true, eps_star, alpha_sigma, sliced_w1, make_dirs, N_EVAL) from ddim_family import uniform_t_grid, make_model, ddim_step _BETAS, _ABAR = linear_schedule() _LAM = 0.5 * np.log(np.maximum(_ABAR, 1e-300) / np.maximum(1.0 - _ABAR, 1e-300)) # ══════════════════════════════════════════════════════════════════ # 0. lambda 坐标与网格 # ══════════════════════════════════════════════════════════════════ def lam_of_abar(abar): """half log-SNR:lambda = log(alpha) - log(sigma) = 0.5 * log(abar/(1-abar))。""" return 0.5 * np.log(np.maximum(abar, 1e-300) / np.maximum(1.0 - abar, 1e-300)) def abar_of_lam(lam): """lambda 的反函数:abar = e^(2 lam) / (1 + e^(2 lam))。""" e = np.exp(2.0 * lam) return e / (1.0 + e) def nearest_t(abar): """把连续的 abar 落到最近的训练时间步上(真实网络只认整数 t)。 注意 abar 是**递减**的,不能用 searchsorted(它只认递增数组), 老老实实找绝对值最小的那个下标。 """ return int(np.clip(int(np.argmin(np.abs(_ABAR - abar))), 1, T)) def uniform_lambda_grid(S: int, quantize: bool = False): """在 lambda 上等间隔取 S 个点,再落回训练网格。 返回 S+1 个 abar:前 S 个是「要评估模型」的位置,最后一个是 1.0 (完全干净,sigma = 0,对应 lambda = +inf,只能用一阶收尾)。 quantize=True 会把每个 lambda 落回整数时间步。落回去之后可能撞车 (两个 lambda 落到同一个 t),撞车会让多步法的差商分母变成 0, 所以强制 t 严格递减。参考解要更细的分辨率,用 quantize=False 直接 吃连续的 abar —— oracle 认得任意 abar,真实网络才需要整数。 """ l0, l1 = _LAM[T], _LAM[1] # 从最吵到最干净 lams = np.linspace(l0, l1, S) if not quantize: return None, np.concatenate([abar_of_lam(lams), [1.0]]) ts = [] for l in lams: t = nearest_t(abar_of_lam(l)) if ts and t >= ts[-1]: t = ts[-1] - 1 # 强制严格递减,避免差商分母为 0 if t < 1: break # 干净端训练网格已无分辨率,到此为止 ts.append(t) ts = np.array(ts, dtype=int) return ts, np.concatenate([_ABAR[ts], [1.0]]) class Counter: """包一层模型,数一下到底调了几次。NFE 不能靠嘴算。""" def __init__(self, model): self.m = model self.n = 0 def __call__(self, x, abar): self.n += 1 return self.m(x, abar) # ══════════════════════════════════════════════════════════════════ # 1. 一阶指数积分器(DPM-Solver-1 / dpmsolver++ 的一阶更新) # ══════════════════════════════════════════════════════════════════ def dpm_pp_first(x, m0, a_s, s_s, a_t, s_t): """x_t = (sigma_t/sigma_s) * x_s - alpha_t * (e^{-h} - 1) * x0_hat h = lambda_t - lambda_s。sigma_t = 0 时 lambda_t = +inf,e^{-h} = 0, 整条式子退化成 x_t = x0_hat —— 最后一步自动正确,不用特判。 """ if s_t <= 0.0: return m0 h = (np.log(a_t) - np.log(s_t)) - (np.log(a_s) - np.log(s_s)) return (s_t / s_s) * x - a_t * (np.exp(-h) - 1.0) * m0 def dpm_pp_second(x, hist, a_s, s_s, a_t, s_t, solver_type="midpoint"): """diffusers multistep_dpm_solver_second_order_update 的逐项复刻。 D1 是用上一步的 x0_hat 做差商估出来的导数:D1 ~ h * dx0/dlambda。 midpoint 和 heun 的区别只在 D1 前面的系数,两者都是二阶。 """ lam_t = np.log(a_t) - np.log(s_t) lam_s0, m0 = hist[-1] lam_s1, m1 = hist[-2] h = lam_t - lam_s0 h0 = lam_s0 - lam_s1 r0 = h0 / h D0, D1 = m0, (m0 - m1) / r0 emh = np.exp(-h) if solver_type == "midpoint": return ((s_t / s_s) * x - a_t * (emh - 1.0) * D0 - 0.5 * a_t * (emh - 1.0) * D1) return ((s_t / s_s) * x - a_t * (emh - 1.0) * D0 + a_t * ((emh - 1.0) / h + 1.0) * D1) def dpm_pp_third(x, hist, a_s, s_s, a_t, s_t, d2_scale=2.0): """三阶多步更新。 d2_scale=1.0 是 diffusers(v0.30.0)`multistep_dpm_solver_third_order_update` 的逐项复刻;d2_scale=2.0 是本文按「二次 x0_hat 必须精确」推出来的系数。 两者的差别见 F 节的受控检验 —— 差的就是这一个 2。 """ lam_t = np.log(a_t) - np.log(s_t) lam_s0, m0 = hist[-1] lam_s1, m1 = hist[-2] lam_s2, m2 = hist[-3] h = lam_t - lam_s0 h0, h1 = lam_s0 - lam_s1, lam_s1 - lam_s2 r0, r1 = h0 / h, h1 / h D0 = m0 D1_0, D1_1 = (m0 - m1) / r0, (m1 - m2) / r1 D1 = D1_0 + (r0 / (r0 + r1)) * (D1_0 - D1_1) D2 = d2_scale * (D1_0 - D1_1) / (r0 + r1) emh = np.exp(-h) return ((s_t / s_s) * x - a_t * (emh - 1.0) * D0 + a_t * ((emh - 1.0) / h + 1.0) * D1 - a_t * ((emh - 1.0 + h) / h ** 2 - 0.5) * D2) # ══════════════════════════════════════════════════════════════════ # 2. 五个采样器 # ══════════════════════════════════════════════════════════════════ def solve_ddim(grid, model, x): """DDIM(eta=0):论文式 12,完全不碰 lambda。""" for i in range(len(grid) - 1): eps = model(x, grid[i]) x, _ = ddim_step(x, grid[i], grid[i + 1], eps, 0.0, 0.0) return x def solve_dpm1(grid, model, x): """一阶指数积分器。理论上应该和 DDIM 逐位相同 —— A 节去验。""" for i in range(len(grid) - 1): a_s, s_s = alpha_sigma(grid[i]) a_t, s_t = alpha_sigma(grid[i + 1]) eps = model(x, grid[i]) m0 = (x - s_s * eps) / a_s x = dpm_pp_first(x, m0, a_s, s_s, a_t, s_t) return x def solve_dpm2s(grid, model, x): """单步二阶(DPM-Solver-2):每步先在 lambda 中点落一脚,用中点的 x0 走完。 一步两步评估。中点的 x0 对积分的「加权平均」是二阶准确的, 所以整步的局部误差是 O(h^3),全局 O(h^2)。 """ for i in range(len(grid) - 1): a_s, s_s = alpha_sigma(grid[i]) a_t, s_t = alpha_sigma(grid[i + 1]) eps = model(x, grid[i]) m0 = (x - s_s * eps) / a_s lam_s, lam_t = np.log(a_s) - np.log(s_s), None if s_t <= 0.0: # 收尾:lambda_t = +inf,中点无从谈起 x = m0 continue lam_t = np.log(a_t) - np.log(s_t) lam_m = 0.5 * (lam_s + lam_t) abar_m = _ABAR[nearest_t(abar_of_lam(lam_m))] a_m, s_m = alpha_sigma(abar_m) x_m = dpm_pp_first(x, m0, a_s, s_s, a_m, s_m) # 先跳到中点 eps_m = model(x_m, abar_m) m_m = (x_m - s_m * eps_m) / a_m # 中点处的 x0 x = dpm_pp_first(x, m_m, a_s, s_s, a_t, s_t) # 用中点 x0 走完整步 return x def _midpoint_x0(x, m0, a_s, s_s, a_t, s_t, model): """在 lambda 的中点补一次评估,拿中点的 x0 走完整步(局部二阶)。 多步法开头那一步没有历史可用,只能降成一阶,而一阶在长度为 h 的区间上 局部误差是 O(h^2),这个误差会一路传到终点 —— 这就是 3M 实测阶数被卡在 2 的原因。给第一步多花一次评估可以验证这件事。 """ lam_s, lam_t = np.log(a_s) - np.log(s_s), np.log(a_t) - np.log(s_t) abar_m = abar_of_lam(0.5 * (lam_s + lam_t)) a_m, s_m = alpha_sigma(abar_m) x_m = dpm_pp_first(x, m0, a_s, s_s, a_m, s_m) eps_m = model(x_m, abar_m) return (x_m - s_m * eps_m) / a_m def solve_multistep(grid, model, x, order=2, solver_type="midpoint", warm2s=False, d2_scale=2.0): """多步法(DPM-Solver++ 2M / 3M):一步一次评估,导数是拿历史 x0 做差商。 开头几步历史不够,自动降阶(和 diffusers 的 lower_order_nums 一样); 最后一步 sigma=0 强制降成一阶(和 final_sigmas_type="zero" 一样)。 warm2s=True 时,第一步改用中点法(多花一次评估)把局部精度提到二阶。 """ hist = [] for i in range(len(grid) - 1): a_s, s_s = alpha_sigma(grid[i]) a_t, s_t = alpha_sigma(grid[i + 1]) eps = model(x, grid[i]) m0 = (x - s_s * eps) / a_s hist.append((np.log(a_s) - np.log(s_s), m0)) last = (i == len(grid) - 2) if i == 0 and warm2s and s_t > 0.0: x = dpm_pp_first(x, _midpoint_x0(x, m0, a_s, s_s, a_t, s_t, model), a_s, s_s, a_t, s_t) elif s_t <= 0.0 or len(hist) < 2 or last: x = dpm_pp_first(x, m0, a_s, s_s, a_t, s_t) elif order >= 3 and len(hist) >= 3: x = dpm_pp_third(x, hist, a_s, s_s, a_t, s_t, d2_scale=d2_scale) else: x = dpm_pp_second(x, hist, a_s, s_s, a_t, s_t, solver_type) return x # ══════════════════════════════════════════════════════════════════ # 3. 主流程 # ══════════════════════════════════════════════════════════════════ def rms_err(x, xref): return float(np.sqrt(((x - xref) ** 2).sum(1).mean())) def fit_slope(ns, errs, res, nfe_min=24): """在双对数上拟合 err ~ N^(-p),返回 p。 只取渐近段:NFE 太小的时候高阶方法还在「攒历史」(前几步被迫降阶), 这一段量出来的斜率既不是 1 也不是 3,什么都不说明;误差小于参考解 分辨率的点是机器精度,同样不能要。 """ xs, ys = [], [] for n, e in zip(ns, errs): if n >= nfe_min and e > 3.0 * res and np.isfinite(e) and e > 0: xs.append(np.log(n)) ys.append(np.log(e)) if len(xs) < 3: return float("nan"), 0 p = -np.polyfit(xs, ys, 1)[0] return float(p), len(xs) def local_err_test(): """受控检验:规定 x0_hat(lambda) 是二次多项式,看各阶更新的局部误差阶。 一条真正的 k 阶方法,在 x0_hat 是 k-1 次多项式时必须一步算准(误差 ~ 0), 因为它的构造就是「用 k 个历史点插值出 k-1 次多项式,再对 e^lambda 精确积分」。 这个检验把「公式对不对」和「轨迹好不好」彻底分开,谁也赖不着谁。 """ rng = np.random.default_rng(3) A, Bc, Cc = rng.standard_normal(2), rng.standard_normal(2), rng.standard_normal(2) g = lambda lam: A + Bc * lam + 0.5 * Cc * lam ** 2 def exact_inc(lam_s, lam_t, n=200001): u = np.linspace(lam_s, lam_t, n) f = np.exp(u)[:, None] * np.stack([g(v) for v in u]) return np.trapezoid(f, u, axis=0) hs = [0.4, 0.2, 0.1, 0.05, 0.025] lam_s = 0.3 out = {} for name in ("一阶", "二阶", "三阶-D2原式", "三阶-D2乘2"): errs = [] for h in hs: lam_t = lam_s + h a_s, s_s = alpha_sigma(abar_of_lam(lam_s)) a_t, s_t = alpha_sigma(abar_of_lam(lam_t)) x = np.array([0.7, -0.4]) hist = [(lam_s - j * h, g(lam_s - j * h)) for j in (2, 1, 0)] # 老→新 if name == "一阶": got = dpm_pp_first(x, g(lam_s), a_s, s_s, a_t, s_t) elif name == "二阶": got = dpm_pp_second(x, hist, a_s, s_s, a_t, s_t) elif name == "三阶-D2原式": got = dpm_pp_third(x, hist, a_s, s_s, a_t, s_t, d2_scale=1.0) else: got = dpm_pp_third(x, hist, a_s, s_s, a_t, s_t, d2_scale=2.0) want = (s_t / s_s) * x + s_t * exact_inc(lam_s, lam_t) errs.append(float(np.abs(got - want).max())) out[name] = errs print(f" {name:>12}: " + " ".join(f"{e:.2e}" for e in errs)) print(" (h 从 0.4 缩到 0.025)") print(" → 一阶 h^2、二阶 h^3 都对;三阶原式只有 h^3(等于二阶),") print(" D2 乘 2 之后直接掉到 1e-14 —— 二次 x0_hat 下它一步就算准了") def main(): rng = np.random.default_rng(20260927) w, mu, Sig = make_target() TH = make_dirs() Xref = sample_true(N_EVAL, w, mu, Sig, rng) model = make_model("oracle", w, mu, Sig) # ── A. DDIM == 一阶 DPM-Solver ? ── print("=" * 72) print("A. 交叉验证:DDIM(eta=0) 与 lambda 空间一阶指数积分器是不是同一个东西") print("=" * 72) n = 500 for S in (20, 50): _, grid_l = uniform_lambda_grid(S) _, grid_t = uniform_t_grid(S) for name, grid in (("uniform-t", grid_t), ("uniform-lambda", grid_l)): x0 = rng.standard_normal((n, 2)) xa = solve_ddim(grid, model, x0.copy()) xb = solve_dpm1(grid, model, x0.copy()) print(f" S={S:3d} {name:>14} max|DDIM - DPM-Solver-1| = " f"{np.abs(xa - xb).max():.3e}") print(" → 两套代码、两套记号,输出差在浮点误差量级:DDIM 就是一阶方法") # ── B. 参考解 ── print() print("=" * 72) print("B. 参考解:3M 跑 4000 步当真值,再自查一下它自己收敛到哪") print("=" * 72) n_ref = 1000 x_init = rng.standard_normal((n_ref, 2)) _, g_fine = uniform_lambda_grid(4000, quantize=False) _, g_mid = uniform_lambda_grid(2000, quantize=False) xref = solve_multistep(g_fine, model, x_init.copy(), order=3) xchk = solve_multistep(g_mid, model, x_init.copy(), order=3) res = rms_err(xref, xchk) print(f" 4000 步 vs 2000 步参考解的差 = {res:.3e}") print(f" → 这个数就是下面所有误差表的「分辨率」,比它小的差别不可信") # ── C. 阶数 ── print() print("=" * 72) print("C. 把阶数量出来:终点到参考解的 RMS 距离 vs 真实 NFE") print("=" * 72) NS = [6, 8, 12, 16, 24, 32, 48, 64, 96, 128, 192] solvers = { "DDIM (uniform-t)": lambda S: (uniform_t_grid(S)[1], solve_ddim, {}), "DDIM (=DPM-1, lambda)": lambda S: (uniform_lambda_grid(S)[1], solve_ddim, {}), "DPM-Solver-2 单步": lambda S: (uniform_lambda_grid(max(S // 2, 2))[1], solve_dpm2s, {}), "DPM-Solver++ 2M": lambda S: (uniform_lambda_grid(S)[1], solve_multistep, {"order": 2}), "3M(D2 用原式)": lambda S: (uniform_lambda_grid(S)[1], solve_multistep, {"order": 3, "d2_scale": 1.0}), "3M(D2 修正×2)": lambda S: (uniform_lambda_grid(S)[1], solve_multistep, {"order": 3, "d2_scale": 2.0}), } print(f" {'NFE':>5} " + " ".join(f"{k:>14}" for k in solvers)) table = {k: [] for k in solvers} nfes = [] for S in NS: row, real_n = [], [] for k, build in solvers.items(): grid, fn, kw = build(S) c = Counter(model) x = fn(grid, c, x_init.copy(), **kw) row.append(rms_err(x, xref)) real_n.append(c.n) nfes.append(int(np.mean(real_n))) table_k = list(solvers) for k, v in zip(table_k, row): table[k].append(v) print(f" {nfes[-1]:5d} " + " ".join(f"{v:14.3e}" for v in row)) print() print(f" 拟合误差 ~ NFE^(-p)(只取 NFE>=24 且误差高于分辨率 {res:.1e}×3 的渐近段):") for k in solvers: p, m = fit_slope(nfes, table[k], res) print(f" {k:>24} p = {p:.2f} (用了 {m} 个点)") # ── D. SW1 走到哪一步就分不出高下了 ── print() print("=" * 72) print("D. 分布层面的尺子(SW1):重复换起点,看它什么时候失灵") print("=" * 72) floors = [sliced_w1(sample_true(N_EVAL, w, mu, Sig, np.random.default_rng(2000 + r)), Xref, TH=TH) for r in range(8)] print(f" 真样本 vs 真样本(5 次):{np.mean(floors):.4f} ± {np.std(floors):.4f}" f" ← 这就是尺子的分辨率") PS = [10, 20, 50] print(f" {'NFE':>5} " + " ".join(f"{k:>14}" for k in solvers)) for S in PS: row = [] for k, build in solvers.items(): grid, fn, kw = build(S) vs = [] for r in range(5): # 换几份起点,量出估计量的抖动 c = Counter(model) x = fn(grid, c, np.random.default_rng(900 + S * 100 + r) .standard_normal((N_EVAL, 2)), **kw) vs.append(sliced_w1(x, Xref, TH=TH)) row.append((np.mean(vs), np.std(vs))) print(f" {S:5d} " + " ".join(f"{v[0]:9.4f}±{v[1]:.3f}" for v in row)) print(" → 20 步时各采样器的差别还大于抖动;到 50 步大家都贴着地板,") print(" SW1 已经回答不了「谁更好」——这种时候只能看 C 节的确定性误差") # ── E. 同一个精度,省多少次评估 ── print() print("=" * 72) print("E. 把终点误差压到 1e-2 / 1e-3,最少要几次模型评估(实测扫描)") print("=" * 72) scan = [4, 6, 8, 10, 12, 16, 20, 24, 32, 40, 48, 64, 80, 96, 128, 160, 192, 256, 384] for k, build in solvers.items(): got = {} for S in scan: if len(got) == 2: break grid, fn, kw = build(S) c = Counter(model) e = rms_err(fn(grid, c, x_init.copy(), **kw), xref) for tol in (1e-2, 1e-3): if tol not in got and e <= tol: got[tol] = c.n s1 = got.get(1e-2, None) s2 = got.get(1e-3, None) f1 = f"NFE={s1:3d}" if s1 else " >400 " f2 = f"NFE={s2:3d}" if s2 else " >400 " print(f" {k:>24} 1e-2: {f1} 1e-3: {f2}") # ── F. 受控检验:三阶更新到底几阶 ── print() print("=" * 72) print("F. 受控检验:规定一条解析的 x0_hat(lambda),直接量局部误差") print("=" * 72) print(" 做法:不让 x0_hat 由 x 决定,而是规定它是 lambda 的二次多项式") print(" g(lambda)。此时精确解就是 (sigma_t/sigma_s) x + sigma_t * 积分 e^lambda g。") print(" x0_hat 是二次的,所以**任何真正的三阶方法都必须一步算准**。") local_err_test() # ── G. 卡住 3M 的不是开头那一步 ── print() print("=" * 72) print("G. 是不是「第一步被迫降阶」拖累的?给第一步多花一次评估试试") print("=" * 72) for tag, kw in (("第一步一阶(默认)", {}), ("第一步改中点法(多 1 次评估)", {"warm2s": True})): ns, es = [], [] for S in (24, 32, 48, 64, 96, 128, 192): grid = uniform_lambda_grid(S)[1] c = Counter(model) ns.append(c.n if False else S) es.append(rms_err(solve_multistep(grid, c, x_init.copy(), order=3, d2_scale=1.0, **kw), xref)) p, _ = fit_slope(ns, es, res, nfe_min=24) print(f" 3M(D2 原式)+ {tag:<22} 实测阶数 p = {p:.2f}") print(" → 没变化。第一步不是瓶颈,瓶颈在系数本身(见 F 节)") if __name__ == "__main__": main() spacing_lab.py # -*- coding: utf-8 -*- """步数该往哪放:同样 20 次评估,换个摆法能差出一个数量级。 阶数(dpm_solver_lab.py 里量出来的 p)说的是「步数变多时误差掉多快」, 但没说「步数摆在哪」。这一篇的主角是**时间步的摆法**: - 在 t 上等间隔(DDIM 论文 / diffusers 默认 leading) - 在 lambda 上等间隔(DPM-Solver 的建议) - 在 lambda 上等间隔但落回整数 t(真实网络只能吃整数) - 在 log sigma 上等间隔 - Karras / EDM 的 rho=7 摆法(SDXL、EDM 系列在用) 判断标准是两条:确定性误差(终点到参考解的距离)和「难度曲线」—— 沿着一条高精度参考轨迹量 |dx0_hat / d lambda|,看哪一段 lambda 上 x0_hat 变得最快。步数就该往那儿放。 /usr/local/bin/python3 spacing_lab.py 依赖:numpy(+ 同目录的 oracle_gmm.py / ddim_family.py / dpm_solver_lab.py) """ import numpy as np from oracle_gmm import (T, linear_schedule, make_target, sample_true, alpha_sigma, sliced_w1, make_dirs, N_EVAL) from ddim_family import uniform_t_grid, make_model, ddim_step from dpm_solver_lab import (uniform_lambda_grid, solve_ddim, solve_multistep, solve_dpm2s, lam_of_abar, abar_of_lam, nearest_t, Counter, rms_err) _BETAS, _ABAR = linear_schedule() _LAM = lam_of_abar(_ABAR) # ══════════════════════════════════════════════════════════════════ # 1. 五种摆法 # ══════════════════════════════════════════════════════════════════ def grid_uniform_t(S: int): """在 t 上等间隔(diffusers 的 leading / DDIM 论文默认)。""" return uniform_t_grid(S)[1] def grid_uniform_lambda(S: int): """在 lambda 上等间隔,连续 abar(DPM-Solver 理论里的标准摆法)。""" return uniform_lambda_grid(S, quantize=False)[1] def grid_uniform_lambda_q(S: int): """在 lambda 上等间隔,但落回整数时间步(真实网络只认整数 t)。""" return uniform_lambda_grid(S, quantize=True)[1] def _ratio(abar): """EDM / diffusers 语境里的 sigma:sigma_t / alpha_t = e^{-lambda}。 注意这不是 sqrt(1-abar)。VP 调度下 lambda 就是 -log 这个量, 所以「在 log sigma 上等间隔」等于「在 lambda 上等间隔」——两者是同一个东西, 表里因此只留一种。 """ return np.sqrt(np.maximum(1.0 - abar, 1e-300) / np.maximum(abar, 1e-300)) def grid_uniform_lambda_capped(S: int, lam_max: float): """在 lambda 上等间隔,但只走到 lam_max 就收尾。 为什么要这个变体:uniform-t 天然把 lambda 的上界截在某个值(NFE=20 时 是 +1.76),剩下的干净端交给最后那个「直接输出 x0_hat」的一阶收尾步。 如果 uniform-t 赢的原因是「截得早」而不是「t 本身特别」,那么把 uniform-lambda 截到同一个上界,两者就该差不多。 """ l0 = _LAM[T] lams = np.linspace(l0, lam_max, S) return np.concatenate([abar_of_lam(lams), [1.0]]) def grid_karras(S: int, rho: float = 7.0): """Karras / EDM 的摆法:在 sigma^(1/rho) 上等间隔,EDM 原文取 rho=7。 sigma 用的是 EDM 的那个 sigma(= sigma_t/alpha_t),VP 下等于 e^{-lambda}, 区间是 [_ratio(abar_T), _ratio(abar_1)] ≈ [158, 0.01]。 """ r_hi, r_lo = _ratio(_ABAR[T]), _ratio(_ABAR[1]) r = (r_hi ** (1.0 / rho) + np.arange(S) / (S - 1) * (r_lo ** (1.0 / rho) - r_hi ** (1.0 / rho))) ** rho return np.concatenate([1.0 / (1.0 + r ** 2), [1.0]]) GRIDS = { "uniform-t": grid_uniform_t, "uniform-lambda": grid_uniform_lambda, "uniform-lambda(整数t)": grid_uniform_lambda_q, "lambda截到1.8": lambda S: grid_uniform_lambda_capped(S, 1.8), "lambda截到3.0": lambda S: grid_uniform_lambda_capped(S, 3.0), "karras(rho=7)": grid_karras, "karras(rho=3)": lambda S: grid_karras(S, rho=3.0), } # ══════════════════════════════════════════════════════════════════ # 2. 难度曲线:沿着参考轨迹看 x0_hat 在哪一段变得最快 # ══════════════════════════════════════════════════════════════════ def difficulty_profile(model, n=400, steps=1500, seed=5): """返回 (lam_mid, weight):一阶方法每一步的误差贡献权重。 推导:第 i 步的局部误差(在 x 的单位里)是 delta_i = sigma_{t_i} * ∫ e^lambda [x0(lam) - x0(lam_i)] d lambda ≈ sigma_{t_i} * e^{lam_i} * (h^2/2) * x0'(lam_i) = alpha_{t_i} * e^{-h} * (h^2/2) * x0'(lam_i) 也就是说局部误差自带一个 alpha_t 因子:噪声端 alpha≈0,轨迹在 lambda 上 几乎不动,走错一点也不打紧;干净端 alpha≈1,同样的 h 会实打实地错。 所以难度权重是 alpha(lambda) * |dx0_hat / d lambda|,不是单纯的导数。 """ x = np.random.default_rng(seed).standard_normal((n, 2)) grid = uniform_lambda_grid(steps, quantize=False)[1] lams, x0s, alphas = [], [], [] for i in range(len(grid) - 1): a_s, s_s = alpha_sigma(grid[i]) a_t, s_t = alpha_sigma(grid[i + 1]) eps = model(x, grid[i]) m0 = (x - s_s * eps) / a_s lams.append(lam_of_abar(grid[i])) alphas.append(a_s) x0s.append(m0.copy()) x = (s_t / s_s) * x - a_t * (np.exp(-(lam_of_abar(grid[i + 1]) - lam_of_abar(grid[i]))) - 1.0) * m0 if s_t <= 0: break lams = np.array(lams) alphas = np.array(alphas) x0s = np.array(x0s) # (steps, n, 2) dx = np.sqrt(((np.diff(x0s, axis=0)) ** 2).sum(-1).mean(-1)) speed = dx / np.abs(np.diff(lams)) lam_mid = 0.5 * (lams[1:] + lams[:-1]) a_mid = 0.5 * (alphas[1:] + alphas[:-1]) return lam_mid, a_mid * speed, speed # ══════════════════════════════════════════════════════════════════ # 3. 主流程 # ══════════════════════════════════════════════════════════════════ def main(): rng = np.random.default_rng(20260927) w, mu, Sig = make_target() model = make_model("oracle", w, mu, Sig) TH = make_dirs() Xref = sample_true(N_EVAL, w, mu, Sig, rng) n_ref = 1000 x_init = rng.standard_normal((n_ref, 2)) g_ref = uniform_lambda_grid(3000, quantize=False)[1] xref = solve_multistep(g_ref, model, x_init.copy(), order=3) g_chk = uniform_lambda_grid(1500, quantize=False)[1] res = rms_err(xref, solve_multistep(g_chk, model, x_init.copy(), order=3)) print("=" * 72) print("A. 参考解分辨率") print("=" * 72) print(f" 3000 步 vs 1500 步 = {res:.2e}(误差表只能信到这个量级)") # ── B. 摆法对比(确定性误差)── print() print("=" * 72) print("B. 同样 NFE,步数摆在哪:终点到参考解的 RMS 距离") print("=" * 72) print(f" {'摆法':>22} " + " ".join(f"{'N=' + str(s):>9}" for s in (10, 20, 50)) + " | " + " ".join(f"{'N=' + str(s):>9}" for s in (10, 20, 50))) print(f" {'':>22} " + " ".join(f"{'DDIM':>9}" for _ in range(3)) + " | " + " ".join(f"{'2M':>9}" for _ in range(3))) for name, build in GRIDS.items(): row_a, row_b = [], [] for S in (10, 20, 50): grid = build(S) row_a.append(rms_err(solve_ddim(grid, model, x_init.copy()), xref)) row_b.append(rms_err(solve_multistep(grid, model, x_init.copy(), order=2), xref)) print(f" {name:>22} " + " ".join(f"{v:9.3e}" for v in row_a) + " | " + " ".join(f"{v:9.3e}" for v in row_b)) # ── C. 摆法对比(分布层面)── print() print("=" * 72) print("C. 同样的比较,换成 SW1(地板约 0.030,重复换起点看抖动)") print("=" * 72) print(f" {'摆法':>22} " + " ".join(f"{'N=' + str(s):>16}" for s in (10, 20))) for name, build in GRIDS.items(): row = [] for S in (10, 20): grid = build(S) vs = [] for r in range(5): x0 = np.random.default_rng(700 + S * 10 + r).standard_normal( (N_EVAL, 2)) vs.append(sliced_w1(solve_multistep(grid, model, x0, order=2), Xref, TH=TH)) row.append((np.mean(vs), np.std(vs))) print(f" {name:>22} " + " ".join(f"{v[0]:.4f}±{v[1]:.3f}" for v in row)) # ── D. 难度曲线 ── print() print("=" * 72) print("D. 难度曲线:x0_hat 在哪一段 lambda 上变化最快") print("=" * 72) lams, weight, raw = difficulty_profile(model) lo = lams < 0 print(f" 只看 |dx0/dlambda|(不带 alpha 因子):" f"噪声端 {raw[lo].mean():.3e} vs 干净端 {raw[~lo].mean():.3e}") print(f" 乘上 alpha_t 之后的**误差权重**: " f"噪声端 {weight[lo].mean():.3e} vs 干净端 {weight[~lo].mean():.3e}" f" ← 干净端重 {weight[~lo].mean() / weight[lo].mean():.1f} 倍") top = np.argsort(weight)[-5:][::-1] print(" 权重最大的 5 段:") for i in top: print(f" lambda = {lams[i]:+6.2f} 权重 = {weight[i]:.3e}") print(" → alpha_t 这个因子把难度整体推向干净端:噪声端轨迹在 lambda 上几乎") print(" 不动(alpha≈0),走错一点也不打紧;干净端才是一步错步步错的地方") # ── E2. 收尾步不是免费的 ── print() print("=" * 72) print("E2. 收尾步的代价:在 lambda_max 处直接输出 x0_hat,到底偏了多少") print("=" * 72) print(" 沿高精度参考轨迹取快照,对每个 lambda 量") print(" || x0_hat(x_lambda, lambda) - 轨迹终点 || —— 这就是提前收尾的偏差。") g_fine = uniform_lambda_grid(3000, quantize=False)[1] x = x_init.copy() snaps = [] for i in range(len(g_fine) - 1): a_s, s_s = alpha_sigma(g_fine[i]) a_t, s_t = alpha_sigma(g_fine[i + 1]) eps = model(x, g_fine[i]) m0 = (x - s_s * eps) / a_s if i % 60 == 0: snaps.append((lam_of_abar(g_fine[i]), float(np.sqrt(((m0 - xref[:n_ref]) ** 2).sum(1).mean())))) x = (s_t / s_s) * x - a_t * (np.exp(-(lam_of_abar(g_fine[i + 1]) - lam_of_abar(g_fine[i]))) - 1.0) * m0 if s_t <= 0: break for lam, d in snaps[::5]: print(f" lambda = {lam:+6.2f} 提前收尾的偏差 = {d:.3e}") print(" → 偏差随 lambda 单调下降:截得越早,这一步白送的误差越大,") print(" 而且它**不随步数下降**(只取决于截在哪个 lambda)") print() print(" 于是 lambda_max 有个最优值(2M, N=20,只改截断点):") for cap in (1.0, 1.5, 2.0, 2.5, 3.0, 3.5, 4.0, 4.61): grid = grid_uniform_lambda_capped(20, cap) e = rms_err(solve_multistep(grid, model, x_init.copy(), order=2), xref) print(f" lambda_max = {cap:4.2f} 终点误差 = {e:.3e}") print(" → 两头都变差:截太早是收尾偏差,截太晚是每步的 h 变大") # ── E. 20 步时各摆法的落脚点 ── print() print("=" * 72) print("E. NFE=20 时各摆法把步子落在 lambda 的哪几个位置") print("=" * 72) for name, build in GRIDS.items(): grid = build(20) lam = np.array([lam_of_abar(g) if g < 1.0 else np.inf for g in grid]) lam_f = lam[np.isfinite(lam)] print(f" {name:>22}: lambda 从 {lam_f[0]:+.2f} 到 {lam_f[-1]:+.2f}," f"相邻间隔 min {np.diff(lam_f).min():.3f} / max {np.diff(lam_f).max():.3f}") if __name__ == "__main__": main() make_figures.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """给「从 DDIM 到高阶采样器」画五张图。 所有数字都是现场重算的:这里 import ddim_family / dpm_solver_lab 里的 采样器,重新跑一遍实验再画,不抄任何手打的表格。改了那边的实现,重跑 这个脚本图就会跟着变,不会出现「图上是旧数字、正文是新数字」。 五张图分别回答: 1. eta=0 和 eta=1 走的是两条什么样的路(同一份随机数,二维真轨迹) 2. eta 该调多大?oracle 模型下扫一遍,欠拟合模型下再扫一遍 3. 为什么换采样器缓存就得重调:模型输入每步变多少 4. 阶数:误差 vs NFE 的双对数斜率 5. 达到 1e-2 / 1e-3 到底要花几次模型评估 只依赖 numpy + matplotlib。跑法:/usr/local/bin/python3 make_figures.py """ from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np import oracle_gmm as OG import ddim_family as DF import dpm_solver_lab as DL ROOT = Path(__file__).resolve().parents[1] OUT = ROOT / "figures" OUT.mkdir(parents=True, exist_ok=True) plt.rcParams["font.sans-serif"] = ["PingFang SC", "Heiti TC", "Arial Unicode MS"] plt.rcParams["axes.unicode_minus"] = False plt.rcParams["figure.dpi"] = 130 RNG_SEED = 20260927 # ══════════════════════════════════════════════════════════════════ # 公共:目标分布、参考解、评估尺子 # ═════════════════════════════════════════════════════════════════= def build_world(): rng = np.random.default_rng(RNG_SEED) w, mu, Sig = OG.make_target() Xref = OG.sample_true(OG.N_EVAL, w, mu, Sig, rng) TH = OG.make_dirs() return w, mu, Sig, Xref, TH def floors(TH, Xref, w, mu, Sig, n=8, base=2000): """尺子的分辨率:真样本对真样本,量 n 次。""" v = [OG.sliced_w1(OG.sample_true(OG.N_EVAL, w, mu, Sig, np.random.default_rng(base + r)), Xref, TH=TH) for r in range(n)] return float(np.mean(v)), float(np.std(v)) # ══════════════════════════════════════════════════════════════════ # 图 1:两条轨迹 # ═════════════════════════════════════════════════════════════════= def fig_traj(w, mu, Sig, Xref): model = DF.make_model("oracle", w, mu, Sig) S = 20 ts, grid = DF.uniform_t_grid(S) n_show = 4 x_init, zs = DF.paired_noise(S, n_show, seed=4321) fig, ax = plt.subplots(figsize=(6.4, 5.4)) ax.scatter(Xref[:, 0], Xref[:, 1], s=3, c="#C9CDD4", alpha=0.45, label="真实样本(目标分布)") styles = {1.0: ("#D2691E", 0.9, 5, r"$\eta=1$(DDPM 祖采样式):每步都重掷噪声"), 0.0: ("#2F6FB3", 1.7, 8, r"$\eta=0$(DDIM,确定性):一条光滑的路")} for eta, (c, lw, ms, lab) in styles.items(): _, (xs, _) = DF.run_ddim(x_init.copy(), grid, model, eta, zs, track=True) P = np.stack(xs, axis=1) # (n, S+1, 2) for k in range(n_show): ax.plot(P[k, :, 0], P[k, :, 1], "-", color=c, lw=lw, alpha=0.85 if eta == 0.0 else 0.55, zorder=4 if eta == 0.0 else 3, label=lab if k == 0 else None) ax.scatter(P[k, 1:-1, 0], P[k, 1:-1, 1], s=ms, color=c, alpha=0.85 if eta == 0.0 else 0.5, zorder=4) ax.scatter(x_init[:, 0], x_init[:, 1], s=110, marker="*", c="#111111", zorder=6, label="起点(两边共用,只有 4 个)") ax.set_xlabel(r"$x_1$") ax.set_ylabel(r"$x_2$") ax.set_title("同一份起点、同一份随机数:N=20 步走出的两条路") ax.legend(loc="upper right", fontsize=8, framealpha=0.9) ax.set_aspect("equal", adjustable="box") fig.tight_layout() p = OUT / "traj_eta.png" fig.savefig(p) plt.close(fig) print(f" [1] {p.name} 轨迹 {n_show} 条 × 2 种 eta") # ══════════════════════════════════════════════════════════════════ # 图 2:eta 扫描(oracle + 欠拟合) # ═════════════════════════════════════════════════════════════════= def fig_eta_sweep(w, mu, Sig, Xref, TH): NS = (10, 20, 50, 100) etas = (0.0, 0.25, 0.5, 0.75, 1.0) flo, flo_sd = floors(TH, Xref, w, mu, Sig) panels = {} for tag, kind in (("oracle(真 score,模型误差=0)", "oracle"), ("欠拟合(模型以为数据是单高斯)", "gauss")): model = DF.make_model(kind, w, mu, Sig) panels[tag] = {e: [] for e in etas} for e in etas: for S in NS: _, grid = DF.uniform_t_grid(S) xi, zs = DF.paired_noise(S, OG.N_EVAL) x, _ = DF.run_ddim(xi, grid, model, e, zs) panels[tag][e].append(OG.sliced_w1(x, Xref, TH=TH)) fig, axes = plt.subplots(1, 2, figsize=(11.2, 4.4), sharey=True) cmap = plt.get_cmap("viridis") for ax, (tag, tab) in zip(axes, panels.items()): for i, e in enumerate(etas): ax.plot(NS, tab[e], "o-", color=cmap(0.15 + 0.7 * i / (len(etas) - 1)), lw=1.6, ms=5, label=r"$\eta=" + f"{e:g}" + r"$") ax.axhline(flo, ls="--", c="#999999", lw=1.2, label=f"噪声地板 {flo:.4f}") ax.set_xscale("log") ax.set_xticks(list(NS)) ax.set_xticklabels([str(s) for s in NS]) ax.set_xlabel("采样步数 N") ax.set_title(tag) ax.grid(alpha=0.25) axes[0].set_ylabel("SW1(越小越好)") axes[0].legend(fontsize=8) axes[1].legend(fontsize=8) fig.suptitle("eta 越小越好,而且两个模型下结论一致", fontsize=12) fig.tight_layout() p = OUT / "eta_sweep.png" fig.savefig(p) plt.close(fig) print(f" [2] {p.name} 地板 {flo:.4f}±{flo_sd:.4f};" f"oracle N=10: eta0={panels['oracle(真 score,模型误差=0)'][0.0][0]:.4f}" f" eta1={panels['oracle(真 score,模型误差=0)'][1.0][0]:.4f}") # ══════════════════════════════════════════════════════════════════ # 图 3:每步「模型输入变了多少」 # ═════════════════════════════════════════════════════════════════= def fig_cache_jitter(w, mu, Sig): model = DF.make_model("oracle", w, mu, Sig) S = 50 _, grid = DF.uniform_t_grid(S) x_init, zs = DF.paired_noise(S, OG.N_EVAL) fig, ax = plt.subplots(figsize=(6.6, 4.4)) cols = {0.0: "#2F6FB3", 0.5: "#7A7A7A", 1.0: "#D2691E"} last10 = {} for eta, c in cols.items(): _, (xs, _) = DF.run_ddim(x_init.copy(), grid, model, eta, zs, track=True) rel = [float(np.sqrt((((xs[i + 1] - xs[i]) ** 2).sum(1)).mean()) / np.sqrt(((xs[i] ** 2).sum(1)).mean())) for i in range(len(xs) - 1)] last10[eta] = float(np.mean(rel[-10:])) ax.plot(range(1, len(rel) + 1), rel, "-", color=c, lw=1.6, label=r"$\eta=" + f"{eta:g}" + r"$") ax.plot([len(rel) - 9, len(rel)], [last10[eta], last10[eta]], lw=3.2, color=c, alpha=0.35) ax.set_xlabel("步序号(越往右越接近干净端)") ax.set_ylabel(r"模型输入相对变化 $\Delta x_i / x_i$(长度之比,无量纲)") ax.set_title("缓存加速盯的就是这条线:eta 越大,每步输入跳得越狠") ax.legend(fontsize=9) ax.grid(alpha=0.25) ratio = last10[1.0] / last10[0.0] ax.annotate(f"最后 10 步:eta=1 是 eta=0 的 {ratio:.1f} 倍", xy=(0.42, 0.86), xycoords="axes fraction", fontsize=9, bbox=dict(boxstyle="round,pad=0.35", fc="#FFF6E5", ec="#D2691E")) fig.tight_layout() p = OUT / "cache_jitter.png" fig.savefig(p) plt.close(fig) print(f" [3] {p.name} 最后10步 eta0={last10[0.0]:.4f} " f"eta0.5={last10[0.5]:.4f} eta1={last10[1.0]:.4f} 倍数 {ratio:.1f}") # ══════════════════════════════════════════════════════════════════ # 图 4 + 图 5:阶数与评估预算 # ═════════════════════════════════════════════════════════════════= def build_reference(w, mu, Sig): """超高分辨率参考解;顺手用 2000 步自查它的收敛精度。 随机数的取法与 dpm_solver_lab.py 的 B 节一模一样(同一个种子、同样的 消耗顺序:先抽一批真样本、再抽初始噪声),所以这里打印的分辨率和那边 是同一个数,图上和正文不会对不上。 """ model = DF.make_model("oracle", w, mu, Sig) rng = np.random.default_rng(RNG_SEED) OG.sample_true(OG.N_EVAL, w, mu, Sig, rng) x_init = rng.standard_normal((1000, 2)) _, gf = DL.uniform_lambda_grid(4000, quantize=False) _, gm = DL.uniform_lambda_grid(2000, quantize=False) xref = DL.solve_multistep(gf, model, x_init.copy(), order=3, d2_scale=2.0) xchk = DL.solve_multistep(gm, model, x_init.copy(), order=3, d2_scale=2.0) return model, x_init, xref, DL.rms_err(xref, xchk) def fig_order_and_budget(w, mu, Sig): model, x_init, xref, res = build_reference(w, mu, Sig) NS = [6, 8, 12, 16, 24, 32, 48, 64, 96, 128, 192] solvers = { "DDIM(uniform-t)": lambda S: (DF.uniform_t_grid(S)[1], DL.solve_ddim, {}), "DDIM(=DPM-1, lambda)": lambda S: (DL.uniform_lambda_grid(S)[1], DL.solve_ddim, {}), "DPM-Solver-2 单步": lambda S: (DL.uniform_lambda_grid(max(S // 2, 2))[1], DL.solve_dpm2s, {}), "DPM-Solver++ 2M": lambda S: (DL.uniform_lambda_grid(S)[1], DL.solve_multistep, {"order": 2}), "3M(三阶项原式)": lambda S: (DL.uniform_lambda_grid(S)[1], DL.solve_multistep, {"order": 3, "d2_scale": 1.0}), "3M(三阶项修正)": lambda S: (DL.uniform_lambda_grid(S)[1], DL.solve_multistep, {"order": 3, "d2_scale": 2.0}), } cols = ["#2F6FB3", "#6FA8DC", "#7A7A7A", "#3E8E41", "#D2691E", "#B03060"] tab = {k: [] for k in solvers} nfes = [] for S in NS: cn = [] for k, build in solvers.items(): grid, fn, kw = build(S) c = DL.Counter(model) tab[k].append(DL.rms_err(fn(grid, c, x_init.copy(), **kw), xref)) cn.append(c.n) nfes.append(int(round(np.mean(cn)))) # ── 图 4:阶数 ── fig, ax = plt.subplots(figsize=(7.0, 5.0)) for (k, errs), c in zip(tab.items(), cols): p, _ = DL.fit_slope(nfes, errs, res) ax.plot(nfes, errs, "o-", color=c, lw=1.5, ms=4.5, label=f"{k} " + r"$p=" + f"{p:.2f}" + r"$") ax.axhline(res, ls="--", c="#BBBBBB", lw=1.2) ax.text(nfes[-1], res * 1.15, f"参考解分辨率 {res:.1e}", fontsize=8, color="#888888", ha="right") ax.set_xscale("log") ax.set_yscale("log") ax.set_xlabel("NFE(真实模型评估次数,数出来的)") ax.set_ylabel("终点到参考解的 RMS 距离") ax.set_title("斜率就是阶数:低阶方法要加一个数量级的步数才能追上") ax.grid(alpha=0.25, which="both") ax.legend(fontsize=8, loc="lower left") fig.tight_layout() p4 = OUT / "order_nfe.png" fig.savefig(p4) plt.close(fig) # ── 图 5:评估预算 ── scan = [4, 6, 8, 10, 12, 16, 20, 24, 32, 40, 48, 64, 80, 96, 128, 160, 192, 256, 384] got = {} for k, build in solvers.items(): row = {} for S in scan: grid, fn, kw = build(S) c = DL.Counter(model) e = DL.rms_err(fn(grid, c, x_init.copy(), **kw), xref) for tol in (1e-2, 1e-3): if tol not in row and e <= tol: row[tol] = c.n got[k] = row fig, ax = plt.subplots(figsize=(7.4, 4.6)) names = list(solvers) xs = np.arange(len(names)) for j, (tol, lab) in enumerate(((1e-2, r"压到 $10^{-2}$ 所需 NFE"), (1e-3, r"压到 $10^{-3}$ 所需 NFE"))): vals, txt = [], [] for k in names: v = got[k].get(tol) vals.append(v if v else 0) txt.append(str(v) if v else ">" + str(scan[-1])) bars = ax.bar(xs + (j - 0.5) * 0.38, vals, width=0.36, color=["#8FB8DE", "#1F4E79"][j], label=lab) for b, t in zip(bars, txt): ax.text(b.get_x() + b.get_width() / 2, b.get_height() + 4, t, ha="center", fontsize=8) ax.set_xticks(xs) ax.set_xticklabels(names, rotation=18, ha="right", fontsize=8) ax.set_ylabel("NFE") ax.set_title("同一个精度,高阶方法省下的评估次数(柱子顶上标的是实测值)") ax.legend(fontsize=9) ax.grid(alpha=0.25, axis="y") fig.tight_layout() p5 = OUT / "nfe_budget.png" fig.savefig(p5) plt.close(fig) print(f" [4] {p4.name} 参考解分辨率 {res:.3e}") print(f" [5] {p5.name} 1e-2 所需 NFE: " + " ".join(f"{k}={got[k].get(1e-2, '>384')}" for k in names)) print(f" 1e-3 所需 NFE: " + " ".join(f"{k}={got[k].get(1e-3, '>384')}" for k in names)) def main(): print("生成配图(数值全部现场重算)") w, mu, Sig, Xref, TH = build_world() fig_traj(w, mu, Sig, Xref) fig_eta_sweep(w, mu, Sig, Xref, TH) fig_cache_jitter(w, mu, Sig) fig_order_and_budget(w, mu, Sig) print("完成,输出目录:", OUT) if __name__ == "__main__": main()
2026年09月27日
3 阅读
0 评论
0 点赞
2026-09-27
AIGC 基本功|DDPM 训练目标与采样流程-DDPM
DDPM 训练目标与采样流程 所属方向:生成范式 | 难度:进阶 | 前置知识:扩散过程的前向与反向推导(前向闭式解、反向后验、以及 ELBO 化简到「预测噪声」的那一步都在那篇推过,这里直接用结论) 关键词:DDPM、L_simple、ELBO 每步权重、采样方差、Algorithm 1/2、fixed_small / fixed_large 01. 为什么需要它 我第一次照着 DDPM 论文的公式老老实实实现训练目标,结果比「偷懒版」更差。 实验是这样的:数据是二维高斯,调度用原文的线性 $T=1000$,模型是一个跨时间步共享参数的小网络(18 个可学参数),三种目标各训一份,然后算生成分布的精确负对数似然。结果是: 训练目标 期望 NLL(nats,越小越好) 等权 $L_{\text{simple}}$ 1.898760 真 ELBO 每步权重 1.922295 差了 0.0235 nats。而 Ho 等人在原文里就明说了:他们把 ELBO 里那一串只跟 $t$ 有关的系数扔掉,直接等权,反而「sample quality 更好」。我当时以为这只是工程上的凑巧,跑完才发现不是——权重的形状和模型容量是绑在一起的。把同一个网络的时间基从 3 项加到 5 项(30 个参数),结论立刻反过来:真 ELBO 权重 1.880952,等权 1.881758,真 ELBO 反超。 第二个坑是采样方差。我一直把 $\tilde\beta_t$ 当成「真实后验方差」。实测把方差换成真正的后验协方差之后,生成分布的期望 NLL 从 1.881015 降到 1.880914,差的 0.000101 nats 全部来自把方差钉死成 $\tilde\beta_t$;生成协方差与真值的比值从 $[0.98946,\ 0.98713]$ 变成 $[1.0,\ 1.0]$。准确地说:$\tilde\beta_t I$ 是已知 $x_0$ 的条件后验协方差;对边缘反向 $q(x_{t-1}\mid x_t)$,它是全协方差分解中的下界。本例用它生成的分布偏窄,不能把数值幅度推广到所有数据。 第三个观察是权重分配。$t\ge2$ 的 ELBO KL 权重跨 120.4 倍,而低噪声端的噪声预测有很高的不可约误差。不可约误差不产生期望梯度;真正被权重重新分配的是可学习的超额误差及随机梯度噪声。后文的容量对照说明“加权在哪种模型上更合适”需要实测,不能把某个时间段断言成无用功。 这三个坑合起来就是这一篇要讲的事:DDPM 的训练目标和采样循环是两件独立的东西,前者管误差权重,后者管转移;在 $L_{simple}$ 下可以分别选择,但 ELBO 权重显式依赖反向方差,并非完全独立。把这两件事分开看,后面所有改进(DDIM、CFG、flow matching)才读得懂。 02. 最小可用理解 三句话: 训练:每一步都是一个「从 $x_t$ 猜刚才加了什么噪声 $\varepsilon$」的回归问题。DDPM 把它做成等权均方误差 $L_{\text{simple}}$,故意不理 ELBO 给每步分配的系数。 权重:本例 KL 权重 $w_t$ 跨 120 倍。等权是改写目标;是否改善生成质量取决于误差分布、优化、容量和评测目标。 采样:反向一步的均值由 $\varepsilon$ 预测器决定,方差可取固定 $\tilde\beta_t$、$\beta_t$ 或学习值;固定方差不进入 $L_{simple}$,却进入 ELBO,学习方差还需要相应训练项。 这张图要看什么:左轴的两条线是真 ELBO 权重(蓝实线 $\sigma^2=\tilde\beta_t$,红虚线 $\sigma^2=\beta_t$),灰色点线是「等权」压平到 1 的位置。右轴绿线是这一时刻噪声里能学出来的比例 $R^2$。两条曲线正好反向:权重在 $t$ 很小的地方冲到 0.6,而那里的 $R^2$ 几乎是 0;权重最低的 $t\approx350$,反而是 $R^2$ 爬到一半的地方。低噪声端权重大、可预测噪声占比小;右端并不遵循这个反向关系。这提示检查容量分配,不足以单独证明等权必然更优。 03. 数学推导 3.1 ELBO 拆成每步的 KL 变分下界写出来是 $$L = \mathbb{E}_q\Big[-\log p_\theta(x_0|x_1) + \sum_{t=2}^{T} D_{\mathrm{KL}}\big(q(x_{t-1}|x_t,x_0)\,\|\,p_\theta(x_{t-1}|x_t)\big) + D_{\mathrm{KL}}\big(q(x_T|x_0)\,\|\,p(x_T)\big)\Big]$$ 最后一项在固定前向过程与先验时没有可学参数;$L_0=-\log p_\theta(x_0|x_1)$ 也是训练目标的一部分。DDPM 为离散像素采用离散化高斯解码似然,连续数据可选连续密度;中间那一长串是下面推导的 KL 项。把它记为 $L_{t-1}$(下标是 $t-1$ 因为它监督的是「从 $t$ 走到 $t-1$」这一步)。 为什么这一项好算?因为 $q$ 和 $p_\theta$ 都是高斯,两个高斯的 KL 有闭式解。而 $q(x_{t-1}|x_t,x_0)$ 在后验那篇已经推过:$q=\mathcal{N}(\tilde\mu_t,\ \tilde\beta_t I)$,其中 $\tilde\beta_t$ 是一个不依赖 $x_t$ 的常数,$\tilde\mu_t$ 是 $x_t$ 与 $x_0$ 的线性组合。在反向方差固定、不由模型学习的前提下,KL 里唯一依赖模型的部分是均值之差: $$L_{t-1} = \mathbb{E}_q\Big[\tfrac{1}{2\sigma_t^{2}}\big\|\tilde\mu_t(x_t,x_0)-\mu_\theta(x_t,t)\big\|^{2}\Big] + C$$ $C$ 是与参数无关的常数,$\sigma_t^2$ 是反向链在这一步选的方差。 3.2 把均值之差换成噪声之差 $\tilde\mu_t$ 和 $\mu_\theta$ 用同一组系数线性组合,前者组合的是真 $x_0$,后者组合的是模型猜的 $\hat x_0$。把 $x_0=\frac{1}{\sqrt{\bar\alpha_t}}(x_t-\sqrt{1-\bar\alpha_t}\,\varepsilon)$ 和 $\hat x_0=\frac{1}{\sqrt{\bar\alpha_t}}(x_t-\sqrt{1-\bar\alpha_t}\,\varepsilon_\theta)$ 代进去,$x_t$ 那一项整整齐齐地消掉,剩下 $$\tilde\mu_t-\mu_\theta = -\frac{\beta_t}{\sqrt{\alpha_t}\sqrt{1-\bar\alpha_t}}\big(\varepsilon-\varepsilon_\theta\big)$$ 注意这里出现的 $\sqrt{\bar\alpha_{t-1}}/\sqrt{\bar\alpha_t}=1/\sqrt{\alpha_t}$——这是整个化简能成立的关键一步,两个系数只差一个 $\sqrt{\alpha_t}$,所以差值是干净的单项。 平方之后: $$L_{t-1} = \mathbb{E}\Big[w_t\big\|\varepsilon-\varepsilon_\theta\big\|^{2}\Big] + C,\qquad w_t=\frac{\beta_t^{2}}{2\sigma_t^{2}\alpha_t(1-\bar\alpha_t)}$$ 这就是那个「只跟 $t$ 有关的系数」。固定方差时由调度和方差选择决定;若学习方差,就不能把相关项都视为与参数无关的常数。 代入两种方差选择,还能再化简一层: $$w_t=\frac{\beta_t}{2\alpha_t(1-\bar\alpha_{t-1})}\ \ (\sigma_t^{2}=\tilde\beta_t),\qquad w_t=\frac{\beta_t}{2\alpha_t(1-\bar\alpha_t)}\ \ (\sigma_t^{2}=\beta_t)$$ 两条式子的分母只差一个下标。附录 ddpm_lab.py 里解析式和化简式两条都算了,最大相对差 3.93×10⁻¹⁶($\tilde\beta$)和 4.13×10⁻¹⁶($\beta$),就是浮点误差级别。 $t=1$ 时 $\tilde\beta_1=0$,上述固定小方差高斯 KL 权重分母为零、分子非零,不能使用。它不是 $0/0$。ELBO 此时对应重建似然 $L_0$,需另行定义;不是“所有最后一步都不准有高斯密度”。本实验用 $\beta_1>0$ 的连续高斯重建权重补这一点,因此表中的 elbo_tilde 是这一连续数据约定,不是逐字复现离散像素 ELBO。 3.3 $w_t$ 长什么样 线性调度 $T=1000$、$\beta$ 从 $10^{-4}$ 到 $0.02$(原文 CIFAR-10 的配置),$\bar\alpha_T=4.0358\times10^{-5}$。实跑出来: $t$ $\bar\alpha_t$ $\beta_t$ $w_t$($\sigma^2=\tilde\beta_t$) $w_t$($\sigma^2=\beta_t$) 相对 $t=500$ 2 9.998e-01 0.00012 5.9967e-01 2.7269e-01 108.87 10 9.981e-01 0.00028 8.6437e-02 7.3717e-02 15.69 50 9.710e-01 0.00108 1.9279e-02 1.8583e-02 3.50 100 8.970e-01 0.00207 1.0267e-02 1.0081e-02 1.86 250 5.241e-01 0.00506 5.3733e-03 5.3432e-03 0.98 500 7.859e-02 0.01004 5.5082e-03 5.5034e-03 1.00 750 3.351e-03 0.01502 7.6506e-03 7.6502e-03 1.39 900 2.752e-04 0.01801 9.1717e-03 9.1716e-03 1.67 1000 4.036e-05 0.02000 1.0205e-02 1.0204e-02 1.85 最后一列是「相对 $t=500$ 的倍数」,等权相当于把这一列全压成 1.00。 形状是两端高、中间平的一个碗。$\tilde\beta$ 版本的最大值是最小值的 120.4 倍;$\beta$ 版本因为分母用的是 $1-\bar\alpha_t$ 而不是 $1-\bar\alpha_{t-1}$,在 $t$ 很小的地方发散得温和些,动态范围小很多。 左端 $1-\bar\alpha_{t-1}$ 很小,使权重大;中段分母增长与线性增大的 $\beta_t$ 互相抵消;右端 $1-\bar\alpha_t$ 趋近 1 后,$w_t\approx\beta_t/(2\alpha_t)$ 随 $\beta_t$ 上升。不能说 $t=250$ 时分母已饱和:表中此时 $1-\bar\alpha_t\approx0.476$。 这个式子不是我拍脑袋的,脚本里用两个高斯的 KL 数值核了一遍:随机取 $x_0$ 和 $\varepsilon$,令 $\varepsilon_\theta=\varepsilon+\delta$($\delta$ 是一个固定的小扰动),比较「数值算的 KL」与「$w_t\|\delta\|^2$」: t KL(数值) w_t*||delta||^2 相对差 2 1.416848e-02 1.416848e-02 2.89e-13 100 1.612712e-03 1.612712e-03 4.17e-14 500 5.035556e-04 5.035556e-04 1.72e-15 1000 2.188401e-03 2.188401e-03 1.61e-14 相对差在 $10^{-13}$ 量级,推导和代码对上了。 3.4 权重最大的地方,恰恰最学不到东西 $w_t$ 只说明「这一步的误差在 ELBO 里值多少钱」,没说明这一步的误差能不能被压下去。真正该看的是贝叶斯地板:给定 $x_t$ 后 $\varepsilon$ 还剩多少不确定性。 数据是 $\mathcal{N}(\mu_0,S_0)$ 时,$x_t$ 的协方差是 $\bar\alpha_t S_0+(1-\bar\alpha_t)I$,条件方差有闭式解,地板就是 $$\mathrm{floor}(t)=\frac{1}{D}\,\mathrm{tr}\Big[\mathrm{Var}(\varepsilon\,|\,x_t)\Big]=\frac{D-(1-\bar\alpha_t)\,\mathrm{tr}\big[(\bar\alpha_tS_0+(1-\bar\alpha_t)I)^{-1}\big]}{D}$$ 定义可学占比 $R^2(t)=1-\mathrm{floor}(t)$($\varepsilon$ 每维先验方差是 1)。实测:$t=1$ 时 $R^2=3.22\times10^{-4}$,$t=500$ 时 0.9616,$t=1000$ 时 1.0000。 在 $t=1$,噪声系数是 $\sqrt{1-\bar\alpha_1}=0.01$,$R^2=3.22\times10^{-4}$ 即约 0.0322%。这说明噪声有很大的条件方差,不说明最优预测毫无用途:生成所需 score 正由小的条件均值决定。$t=2$ 的 108.87 倍权重与 $t=1$ 的 $R^2$ 不能混为同一个时间步。 因此应把不可约地板与超额误差分开分析。不可约项对参数的期望梯度为零,有限容量如何分配精度、随机梯度方差多大,才决定加权目标的优化表现。 3.5 采样方差不是后验方差 $\tilde\beta_t$ 是 $q(x_{t-1}|x_t,x_0)$ 的方差,条件是「已知 $x_0$」。但采样时我们没有 $x_0$,只有模型猜的 $\hat x_0$,它自己有误差。把这一层不确定性算进去,真实后验协方差是 $$\Sigma_t = \tilde\beta_t I + \Big(\tfrac{\sqrt{\bar\alpha_{t-1}}\beta_t}{1-\bar\alpha_t}\Big)^{2}\mathrm{Var}(x_0-\hat x_0\,|\,x_t)$$ 多出来的第二项就是「均值估计不准」带来的。实测:用最优 $\varepsilon$ 预测器配 $\tilde\beta_t$,生成协方差只有真值的 $[0.98946,\ 0.98713]$;换成真实后验协方差,比值精确回到 $[1.0,\ 1.0]$,NLL 差 $3.56\times10^{-10}$ nats——在本例中极接近数据分布;仍有有限终点先验与数值误差。 这一节的结论对后面很重要:反向链的均值和方差是两件事。均值决定你往哪走,方差决定你抖多厉害;在这个高斯实验中,固定 $\tilde\beta_t$ 使协方差特征值约偏小 1.1%~1.3%;该数值不是通用结论。 04. 代码实现 4.1 调度、$\tilde\beta_t$、$w_t$(20 行) def linear_beta(T=1000, b1=1e-4, bT=0.02): """DDPM 原文 CIFAR-10 用的线性调度。""" return np.linspace(b1, bT, T) class Sched: """下标约定:abar[t]、beta[t]、bt[t] 里的 t 取 1..T;t=0 是数据本身。""" def __init__(self, beta): self.T = len(beta) self.beta = beta self.alpha = 1.0 - beta self.abar = np.concatenate([[1.0], np.cumprod(self.alpha)]) # 长度 T+1 self.bt = np.empty(self.T + 1) self.bt[0] = 0.0 self.bt[1:] = (1.0 - self.abar[:-1]) / (1.0 - self.abar[1:]) * beta def weight(self, which="tilde"): s2 = self.bt[1:] if which == "tilde" else self.beta with np.errstate(divide="ignore", invalid="ignore"): w = self.beta ** 2 / (2.0 * s2 * self.alpha * (1.0 - self.abar[1:])) if which == "tilde": # t=1 时 tilde_beta_1=0,分母为零。DDPM 把这一步交给 L_0 单独处理, # 这里退化地用 fixed_large 的权重顶上,只为让训练能跑。 w[0] = self.beta[0] / (2.0 * self.alpha[0] * (1.0 - self.abar[1])) return w np.errstate 那一行不是装饰:\tilde\beta_1=0 会触发除零,不包起来的话 w/w.mean() 之后整条权重变成 nan,训练静悄悄地训出一个废模型。 4.2 三种目标:共享参数的带权最小二乘 要让「权重决定容量往哪搬」这件事看得见,模型必须跨时间步共享参数——否则每一步各自拟合,权重只影响每一步自己的收敛速度,看不出搬移效果。所以用一个仿射 $\varepsilon$ 预测器加固定时间基: $$\varepsilon_\theta(x_t,t)=W\,\big[x_t\otimes\psi(t),\;\psi(t)\big],\qquad \psi(t)=\big[1,\ (t/T)^{0.5},\ t/T\big]$$ (指数取 $[0,0.5,1]$ 而不是 $[0,1,2]$:因为 $\sqrt{1-\bar\alpha_t}$ 在 $t$ 小的时候像 $\sqrt{t}$,纯整数次幂的多项式基在这一段拟合不出来,曲线会剧烈振荡。这是踩过的坑。) 训练集 $M=120000$ 条 $(x_0,t,\varepsilon)$,特征维度 9,可学参数 18 个。三种权重都归一化到均值 1(整体缩放不改变解,只为了让正则强度可比)。用带权岭回归一步算出闭式解,不做迭代。 4.3 精确 NLL:把蒙特卡洛噪声干掉 一开始我用采样算 NLL,三种目标差 0.002~0.005 nats,而 2 万样本的蒙特卡洛误差就有 ±0.007——结论完全淹没在噪声里。改走闭式解:仿射模型下 $p_\theta(x_0)$ 仍是高斯,把密度沿反向链往前传,最后 $$\mathbb{E}\big[-\log p_\theta(x_0)\big]=H\big[\mathcal{N}(\mu_0,S_0)\big]+D_{\mathrm{KL}}\big(\mathcal{N}(\mu_0,S_0)\,\|\,\mathcal{N}(m,S)\big)$$ 用 40 万样本复核,闭式解 1.880914 对蒙特卡洛 1.880890,差 $2.38\times10^{-5}$。两种独立实现(密度前向传播 / 线性映射复合)的均值差 $9.99\times10^{-16}$、协方差差 $5.00\times10^{-16}$。 三种目标 × 两种采样方差: 训练目标 $\sigma^2=\tilde\beta_t$ $\sigma^2=\beta_t$ 差 生成 std / 真 std uniform(等权) 1.898760 1.895815 −0.002945 0.9659 elbo_tilde 1.922295 1.924692 +0.002396 1.0392 elbo_large 1.914947 1.917129 +0.002181 1.0322 (最后一列是生成分布标准差与真分布标准差的比值,1 表示胖瘦刚好对上。等权明显偏窄,真 ELBO 权重的两个版本反而偏宽。) 4.4 权重把超额误差搬到了哪 「超额」= 实际 MSE − 贝叶斯地板,包括容量、有限样本、正则化与优化误差: $t$ 可学占比 $R^2$ 地板 等权超额 真 ELBO 超额 1 0.0003 0.9997 1.516e-01 1.017e-02 5 0.0022 0.9978 1.120e-01 6.824e-03 20 0.0182 0.9818 5.295e-02 3.897e-03 50 0.0853 0.9147 9.778e-03 7.323e-03 100 0.2510 0.7490 4.271e-03 2.560e-02 200 0.5663 0.4337 1.324e-02 2.820e-02 400 0.9000 0.1000 1.067e-03 4.463e-03 600 0.9876 0.0124 5.375e-03 1.349e-02 800 0.9993 0.0007 3.118e-04 1.082e-03 1000 1.0000 0.0000 1.337e-02 2.345e-02 分档汇总:低噪声档 $t\le50$,等权超额 5.140e-02,真 ELBO 5.137e-03(比值 0.10,好 10 倍);高噪声档 $t\ge400$,等权 3.459e-03,真 ELBO 8.060e-03(比值 2.33,差一倍多)。 这张图要看什么:(a) 两条曲线的交叉点大约在 $t=70$——交叉点左边真 ELBO 权重更准,右边等权更准;(b) 取比值后看得更清楚,灰色填充区(比值 <1)是真 ELBO 占优的低噪声段,红色填充区(比值 >1)是等权占优的中高噪声段。加权不是全面变好,是把误差从一段搬到另一段。等权之所以赢,是因为搬走的那一头($t\le50$)本来误差就大到没救(1.5e-01 对地板 0.9997),而搬来的那一头($t\ge400$)绝对误差只有 1e-03 量级,赔得起。 4.5 换个容量档,结论反过来 时间基从 3 项加到 5 项($[0,0.5,1,1.5,2]$,30 个参数),其他完全不动: 容量档 时间基项数 可学参数 NLL(等权) NLL(真 ELBO) 谁更好 低噪声档超额比 高噪声档超额比 loose 5 30 1.881758 1.880952 真 ELBO 0.71 3.91 tight 3 18 1.898760 1.922295 等权 0.10 2.33 这张表说明两档容量的排序不同,不能推出“容量足够时 ELBO 自然赢”。ELBO、生成 NLL 和感知质量也不是同一个目标;EDM、Min-SNR 同时涉及预条件、噪声分布或多任务梯度冲突,不应仅归因为模型变大。 4.6 真网络版:Algorithm 1 与 Algorithm 2 闭式解毕竟是玩具。附录 ddpm_train.py 给了一份 numpy 手写版:两层 MLP(输入 $10=2+8$ 维时间嵌入,隐层 64,输出 2),手写反向传播 + Adam,batch 256,20000 步,数据是 7 个高斯的混合(一个中心 + 六个环)。 梯度先对:手写反向传播 vs 有限差分,最大相对差 9.92×10⁻¹⁰。Algorithm 1(训练) 照抄原文,唯一变量就是 loss 前面的 $w_t$: def train(sc, w_full, steps, batch, rng, lr=1e-3): p = init_params(rng) opt = Adam(p, lr=lr) for _ in range(steps): # 1: t ~ Uniform({1, ..., T}) 2: x_0 ~ q(x_0) 3: eps ~ N(0, I) t_idx = rng.integers(1, T + 1, size=batch) # (256,) x0 = sample_mix(batch, rng) # (256, 2) eps = rng.standard_normal((batch, D)) # (256, 2) # 4: 一步加噪(闭式解,不需要真的走 t 步) a = sc.abar[t_idx] # (256,) xt = np.sqrt(a)[:, None] * x0 + np.sqrt(1.0 - a)[:, None] * eps # 5: 梯度下降一步 Z = np.concatenate([xt, time_embed(t_idx)], axis=1) # (256, 10) h1, h2, out = forward(p, Z) # (256,64) (256,64) (256,2) opt.step(p, backward(p, Z, h1, h2, out, eps, w_full[t_idx - 1])) return p 三个 shape 值得停下来看一眼。第一,t_idx 是一个长度为 batch 的随机整数向量,不是标量——每个样本走不同的时间步,这是 $L_{\text{simple}}$ 里那个 $\mathbb{E}_t$ 的实现方式,也是「等权」的字面含义:每个 $t$ 被抽中的概率相同。第二,Z 的第二维是 $10 = 2 + 8$,其中 8 维来自 TIME_FREQ = (1,2,4,8) 的 cos/sin 时间嵌入;不把 $t$ 喂进去的话,网络不知道当前噪声档位,$L_{\text{simple}}$ 根本学不动。第三,w_full[t_idx - 1] 这个减一是全文最容易写错的地方:时间步 $t$ 从 1 数到 $T$,而数组下标从 0 开始。写反了整条权重会整体错位一步,$t$ 很小的地方拿到的是 $t+1$ 的权重——在动态范围 120 倍的曲线上,错位一步就能让训练目标面目全非。 Algorithm 2(采样) 是完整 1000 步链: def sample_chain(sc, p, n, which="tilde", rng=None): x = rng.standard_normal((n, D)) # x_T ~ N(0, I) for t in range(T, 0, -1): eps_hat = predict_eps(p, x, t) a, a_prev = sc.abar[t], sc.abar[t - 1] alpha_t, beta_t = sc.alpha[t - 1], sc.beta[t - 1] mu = (x - beta_t / np.sqrt(1.0 - a) * eps_hat) / np.sqrt(alpha_t) sigma = np.sqrt(sc.bt[t]) if which == "tilde" else np.sqrt(beta_t) z = rng.standard_normal((n, D)) x = mu + (sigma * z if t > 1 else 0.0) # t == 1 时不加噪声 return x 最后那行 if t > 1 else 0.0 就是 diffusers 里 if t > 0 的同一个判断(下标约定差 1):最后一步不加噪声。均值那一行是 $\tilde\mu_t$ 的另一种写法——把 $x_0$ 用 Tweedie 公式替掉之后 $\tilde\mu_t=\frac{1}{\sqrt{\alpha_t}}\big(x_t-\frac{\beta_t}{\sqrt{1-\bar\alpha_t}}\varepsilon_\theta\big)$,比 3.2 节那个两项组合的形式更紧凑,也是 diffusers 之外大多数实现采用的写法。 采样之后用三个指标打分:$\log p_0$ 相对真样本、到最近模式中心的距离比、模式覆盖数(某个分量占比超过 2% 才算被覆盖)。这三个指标是互补的——$\log p_0$ 只衡量落在高密度区,模式覆盖才告诉你有没有整个丢掉某个模式;实测两个网络都是 7/7,说明等权赢的不是覆盖面,是落在模式里的精度。 训练目标 $\log p_0$ 相对真样本 到最近中心距离 / 真样本 模式覆盖 uniform −0.2208 ± 0.0195 1.0857 7/7 elbo_tilde −0.4212 ± 0.0062 1.1997 7/7 多峰数据上等权依然赢,而且差距比高斯那组更明显(0.20 vs 0.42 nats)。每步 MSE 剖面也和闭式解的预测一致:$t=1$ 处两个网络都是 ≈1.0(学不动),$t=1000$ 处 uniform 0.0016、elbo_tilde 0.0026(都学得很好)。 只换采样方差、不动网络: 方差选择 $\log p_0$ 相对真样本 到最近中心距离 / 真样本 $\tilde\beta_t$(fixed_small) −0.2363 ± 0.0123 1.0968 $\beta_t$(fixed_large) −0.2640 ± 0.0129 1.1099 这里和 4.3 的高斯实验结论相反:高斯闭式解里 $\beta_t$ 把偏窄补回来了(std 比 0.9659 → 0.9720),多峰数据上 $\tilde\beta_t$ 反而更好。原因不神秘——两者的相对差主要集中在低噪声段,高噪声端反而很接近(见 3.3 节),多峰分布的模式之间本来就脆弱,注入更多噪声会把样本推离模式中心(距离比 1.0968 → 1.1099 就是这个效果)。方差选择取决于数据、预测误差、调度与评价指标,不能只按是否多峰判断,别把一维高斯的结论直接搬。 05. 工业级实现对照 对照 huggingface/diffusers 的 src/diffusers/schedulers/scheduling_ddpm.py → DDPMScheduler.step(以 2026-09 时的实现为准)。这一份源码就是知识树里给这个节点配的 code_refs。 第一处差异:方差有六个分支,不是一个数。 我 4.1 节只写了 $\tilde\beta_t$ 和 $\beta_t$ 两种,框架里是 _get_variance() 的六个 variance_type: if variance_type == "fixed_small": # 默认:后验方差的下界 variance = variance elif variance_type == "fixed_small_log": # 取 log 再 exp(0.5*log),代数上等价 variance = torch.log(variance); variance = torch.exp(0.5 * variance) elif variance_type == "fixed_large": # 直接上 beta_t variance = current_beta_t elif variance_type == "fixed_large_log": variance = torch.log(current_beta_t) elif variance_type == "learned": # 网络自己输出 return predicted_variance elif variance_type == "learned_range": # 在 [tilde_beta, beta] 之间插值 min_log = torch.log(variance); max_log = torch.log(current_beta_t) frac = (predicted_variance + 1) / 2 variance = frac * max_log + (1 - frac) * min_log 值得注意的是 learned_range:它插值区间的两个端点恰好就是我 4.5 节扫描的那两个,而且是在 log 空间插值(这更合理,因为这两个量跨好几个数量级)。也就是说「$\tilde\beta_t$ 还是 $\beta_t$」在工业实现里不是二选一,而是交给网络学一个位置。 第二处差异:最后一步不加噪声。 源码里 variance = 0; if t > 0: ...。这正对应 3.2 节说的 $L_0$ 退化问题——$t=1$ 时 $\tilde\beta_1=0$,这是该采样算法的终步约定;连续数据模型仍可选择非零方差的终步重建分布。 第三处差异:累计噪声表预先计算,转移系数在 step 内由当前与目标时间步组合。 step() 里那两行: pred_original_sample_coeff = (alpha_prod_t_prev ** 0.5 * current_beta_t) / beta_prod_t current_sample_coeff = current_alpha_t ** 0.5 * beta_prod_t_prev / beta_prod_t pred_prev_sample = pred_original_sample_coeff * pred_original_sample + current_sample_coeff * sample 展开就是 3.2 节的 $\tilde\mu_t$,$c_0=\frac{\sqrt{\bar\alpha_{t-1}}\beta_t}{1-\bar\alpha_t}$、$c_x=\frac{\sqrt{\alpha_t}(1-\bar\alpha_{t-1})}{1-\bar\alpha_t}$,一字不差。 这张图要看什么:三条线是反向一步的三个配料随 $t$ 的变化——蓝线是「拉向 $\hat x_0$ 的系数」$c_0$,灰虚线是「保留 $x_t$ 的系数」$c_x$,红线是注入噪声的标准差 $\sqrt{\tilde\beta_t}$。绝大多数步子里 $c_0$ 在 1e-4 ~ 1e-2 量级,也就是一步只挪千分之几;只有最后几十步 $c_0$ 才冲到 1 附近,真正开始「成形」。这是完整 1000 步离散表下的系数,不意味着采样必须走满 1000 步。重排噪声表或换求解器后,模型可用更少的评估次数采样。 第四处差异:prediction_type 有三个选项。 epsilon(默认)、sample(直接预测 $x_0$)、v_prediction(预测 $v$)。源码把三者统一先转成 pred_original_sample,再走同一套系数——参数化只是换了个入口,反向一步的代数完全一样。参数化之间的权重差异在那篇《扩散过程的前向与反向推导》里算过($\varepsilon$ 1.0 倍、$v$ 2.5×10⁴、$x_0$ 2.5×10⁸)。 第五处:clip_sample=True 是默认开的。 把 pred_original_sample 夹到 $[-1,1]$。这在像素空间合理,在潜空间是错的——SD1.x 的 VAE 原始潜变量乘约 0.18215 后才送入扩散网络,尺度约归一到 1;但它并不被限制在 $[-1,1]$。因此不能照搬像素裁剪,是否关闭或采用其他限幅应遵循具体模型配置。 06. 代价与边界 等权省了什么:不用管 $\tilde\beta_1=0$ 的退化,不用管 $L_0$ 的离散解码器,没有显式时间权重,但不保证每步的梯度范数一致(这点很重要——那篇前向反向推导里算过,$x_0$ 参数化跨 2.5×10⁸ 倍,梯度会被最吵的几步吃光)。 等权赔了什么:三档实测摆在一起看—— 场景 等权 真 ELBO 权重 谁赢 容量宽裕(30 参数,高斯) 1.881758 1.880952 真 ELBO 容量吃紧(18 参数,高斯) 1.898760 1.922295 等权 容量吃紧(MLP,7 峰混合) −0.2208 −0.4212 等权 怎么选择权重:若关注低噪声重建精度,应在任务数据上测量各时段超额误差与最终质量,比较等权与加权;本例的 10 倍误差比不能直接推出超分、修复都不该用等权。 采样方差的边界:高斯数据上 $\beta_t$ 把偏窄从 3.4% 补到 2.8%(补回 0.6 个百分点,还没补满);多峰数据上 $\beta_t$ 反而更差。所以别把「fixed_large 更大更对」当成通例——要不要更大方差,取决于你的数据是不是多峰、模式之间经不经得起抖。 这张图要看什么:(a) 三条线是本高斯实验每步的方差——$\tilde\beta_t$(蓝)、$\beta_t$(红虚)、真实后验方差(绿)。绿线与蓝线的差异取决于时间步,较低噪声端尤其需关注,这个差值就是 3.5 节说的「均值估计不准」那一项,也是 NLL 上那 0.0001 nats 的全部来源。(b) 在 $\tilde\beta_t$ 与 $\beta_t$ 之间插值扫一遍:蓝线(左轴)是生成分布标准差与真分布之比,从 0.9659 单调爬到 0.9720;红线(右轴)是期望 NLL,几乎是一条平线。标准差之比是敏感指标,NLL 对这件事几乎不动——想判断方差选得对不对,同时检查 NLL 与协方差,不要只看一个指标。 采样成本要按同一口径计数:4000 样本 × 1000 步是 $4\times10^6$ 次样本级前向;训练 20000 步 × batch 256 是 $5.12\times10^6$ 次样本级前向,另有反向计算。不能拿样本级采样前向数除以训练优化器步数,声称采样贵 200 倍。批量大小与硬件利用率还会改变墙钟时间。 最后一个容易搞混的地方:训练目标(等权还是加权)和采样器(DDPM 随机链还是 DDIM 确定链)是两个正交的旋钮。本篇从头到尾只动前一个,采样器始终是 DDPM 原文那条随机链。换了训练目标不影响你能不能换采样器,反过来也一样——DDIM 那篇最关键的观察就是:DDPM 的训练目标根本没有约束反向链必须是一阶 Markov 的。把这两件事当成一件事,是读扩散模型文献时最普遍的混线。 07. 经典论文脉络 arXiv:1503.03585(Sohl-Dickstein et al., 2015)——扩散的雏形。 定义了前向 Markov 链和反向链,用前向扩散与变分目标训练反向过程,样本质量也远不够看。贡献是「这个方向存在」。 arXiv:2006.11239(Ho et al., 2020)——DDPM,本篇的锚点。 三件事:把 $\tilde\mu_t$ 参数化成预测 $\varepsilon$;指出 ELBO 里那串只跟 $t$ 有关的系数可以扔掉、等权反而更好;给出 $\tilde\beta_t/\tilde\mu_t$ 的闭式解并配上线性调度。这才是「扩散模型能训起来」的直接原因。 arXiv:2102.09672(Nichol & Dhariwal, 2021)——Improved DDPM。 两件事直接对着本篇的洞:一是学方差(对应 learned_range,在 $\tilde\beta_t$ 与 $\beta_t$ 之间让网络选位置,且在 log 空间插值);二是余弦调度,改善低分辨率设置下线性噪声调度过快破坏信号、后段接近纯噪声的问题。 arXiv:2010.02502(Song et al., 2020)——DDIM。 指出 DDPM 的训练目标其实没有约束反向链必须是 Markov 的,于是可以推一个非 Markov 的确定性问题,一步跨很多步。训练和采样在这里正式解耦。 arXiv:2206.00364(Karras et al., 2022)——EDM。 从「每步该加权多少」重新出发,把加权、预处理、调度统一成一套设计空间,并用二阶求解器把步数压到几十步。这是 4.5 节那个「容量变了结论会变」在真实模型上的落地版本。 08. 常见误解 ①「$L_{\text{simple}}$ 就是 ELBO。」 不是。它丢掉了两样东西:每步的 $w_t$,以及 $L_0$ 那一项的离散解码器。它是一个设计选择,不是近似——丢掉的部分在数学上并不小($w_t$ 跨 120 倍),只是恰好在容量吃紧时更划算。 ②「$\tilde\beta_t$ 是真实后验方差。」 不是,它是已知 $x_0$ 时的后验方差。采样时 $x_0$ 是猜的,猜错的那部分不确定性没算进去。实测真实后验方差在中后段比 $\tilde\beta_t$ 高一大截(图 3a 的绿线与蓝线),代价是生成分布偏窄 1.3%。 ③「权重大意味着更容易学。」 权重衡量误差在目标中的代价,条件方差衡量不可约误差;两者不是同一个量。低噪声端 $R^2$ 小不意味着该步 score 不重要。 ④「采样方差是训出来的。」 默认不是。fixed_small 是硬编码的 $\tilde\beta_t$,跟训练毫无关系;只有 learned / learned_range 才让网络参与,而且此时模型输出通道要翻倍(step() 里 model_output.shape[1] == sample.shape[1] * 2 那个判断就是干这个的)。 ⑤「换方差只影响采样,不影响训练。」 一半对。方差确实不参与 $L_{\text{simple}}$ 的计算,但 $w_t$ 的公式里有 $\sigma_t^2$——所以用真 ELBO 权重训练时,你选的方差会通过 $w_t$ 反过来改变训练目标(表 3 里 elbo_tilde 和 elbo_large 是两种不同的训练目标,不只是两种采样方式)。 ⑥「等权赢,所以加权分析没用了。」 本文两个容量档已经出现排序翻转;权重还与噪声调度、预条件、数据与优化相互作用。应在同一评测口径下选择,不能把一次玩具胜负当成普遍结论。 ⑦「$T=1000$ 是个需要调的超参。」 不完全是。$T$ 决定了 $w_t$ 的动态范围,也决定了每步要挪多远;但它同时被 $\beta$ 调度绑住——改 $T$ 不改 $\beta$ 的端点,等于改了整条噪声表的形状。Improved DDPM 换余弦调度而不是换 $T$,正是因为这两个量不能分开调。 09. 动手验证 两个脚本都在附录,只依赖 numpy: /usr/local/bin/python3 ddpm_lab.py # 闭式解实验室,几秒跑完 /usr/local/bin/python3 ddpm_train.py --steps 20000 --n-gen 4000 # 真网络版,几分钟 ddpm_lab.py 会依次打印六段(真实输出摘要,不是预期值): A 段:两种方差下的 $w_t$ 表格,解析式与化简式的最大相对差 3.93e-16;用两个高斯的 KL 数值核对,相对差 2.89e-13;并明确打印 t=1 的 tilde_beta_1 = 0.000e+00。 B 段:训练集规模(120000 条、时间基 3 项、18 个参数)与权重动态范围 120.4。 C 段:真分布熵 1.880914 nats;「最优 $\varepsilon$ 预测器 + $\tilde\beta_t$」1.881015(差 +0.000101);「最优 $\varepsilon$ 预测器 + 真实后验方差」1.880914(差 +3.56e-10)。三种目标 × 两种方差的 NLL 表。 D 段:每步超额误差剖面,低噪声档比值 0.10、高噪声档比值 2.33。 E 段:$\lambda$ 插值扫描,NLL 1.898760 → 1.895815,生成 std 比 0.9659 → 0.9720。 F 段:两档容量对照,loose 真 ELBO 赢(1.880952 vs 1.881758)、tight 等权赢(1.898760 vs 1.922295)。 ddpm_train.py 的 A 段是梯度检查(有限差分相对差 9.92e-10),C 段是三种指标打分(uniform −0.2208±0.0195 / elbo_tilde −0.4212±0.0062),D 段只换采样方差($\tilde\beta_t$ −0.2363±0.0123 / $\beta_t$ −0.2640±0.0129)。 想自己动手改的话,三个地方最值得试:把 BASIS 从 3 项换成 5 项看结论翻转(F 段已经在做);把 T 从 1000 降到 100 看 $w_t$ 的动态范围怎么变;把数据从高斯换成 ddpm_train.py 的 7 峰混合,看方差选择的结论会不会反过来。 10. 延伸阅读 前置:扩散过程的前向与反向推导(前向闭式解、反向后验、ELBO→$\varepsilon$ 的化简)、变分下界与重参数化(ELBO 这个工具本身从哪来)。 已发布后继:DDIM 与高阶采样器、CFG、潜空间扩散、流匹配。 往上一层:把 RL 用到扩散模型上(训练目标被换成奖励之后,$w_t$ 这套分析还成不成立)。 附录:完整代码 09 节用到的脚本全文如下(ddpm_lab.py、ddpm_train.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 ddpm_lab.py # -*- coding: utf-8 -*- """DDPM 训练目标的最小实验室(闭式解版,只依赖 numpy,无 torch)。 这个脚本回答四个只能用数字回答的问题: A. ELBO 里那一串「只跟 t 有关的系数」到底是什么形状。 解析式 w_t = beta_t^2 / (2 sigma_t^2 alpha_t (1 - abar_t)), 用两个高斯的 KL 数值核一遍,确认没推导错。 B. 用带权最小二乘训练一个**跨时间步共享参数**的 epsilon 网络 (仿射 + 固定时间基),比较三种目标: uniform —— DDPM 的 L_simple,等权 elbo_tilde —— 真 ELBO 权重,sigma^2 = tilde_beta_t(diffusers 的 fixed_small) elbo_large —— 真 ELBO 权重,sigma^2 = beta_t(diffusers 的 fixed_large) 模型容量故意给小(9 个时间基函数),所以权重真的会决定容量往哪搬。 C. 仿射模型下 p_theta(x_0) 是一个高斯,可以算出**精确 NLL**。 用「密度沿链前向传播」算,再用「线性映射复合」交叉验证一遍。 D. 反向链的方差 sigma^2 到底该取 tilde_beta_t 还是 beta_t。 在 beta_tilde 与 beta 之间插值扫一遍,看 NLL 与生成分布的胖瘦怎么变。 运行: /usr/local/bin/python3 ddpm_lab.py 依赖: numpy(无 GPU、无 torch) """ import os import numpy as np RNG_SEED = 20260927 T = 1000 D = 2 # ══════════════════════════════════════════════════════════════════ # 0. 噪声调度与前向过程的闭式量 # ══════════════════════════════════════════════════════════════════ def linear_beta(T=1000, b1=1e-4, bT=0.02): """DDPM 原文 CIFAR-10 用的线性调度。""" return np.linspace(b1, bT, T) def cosine_beta(T=1000, s=0.008, clip=0.999): """Improved DDPM 的余弦调度:先定 abar(t) 再反解 beta。""" u = np.arange(1, T + 1) / T f = np.cos(((u + s) / (1 + s)) * np.pi / 2) ** 2 f0 = np.cos((s / (1 + s)) * np.pi / 2) ** 2 abar = f / f0 beta = 1.0 - abar / np.concatenate([[1.0], abar[:-1]]) return np.clip(beta, 1e-8, clip) class Sched: """下标约定:abar[t]、beta[t]、bt[t] 中的 t 取 1..T;t=0 是数据本身。""" def __init__(self, beta): self.T = len(beta) self.beta = beta self.alpha = 1.0 - beta self.abar = np.concatenate([[1.0], np.cumprod(self.alpha)]) # 长度 T+1 # tilde_beta_t = (1 - abar_{t-1}) / (1 - abar_t) * beta_t self.bt = np.empty(self.T + 1) self.bt[0] = 0.0 self.bt[1:] = (1.0 - self.abar[:-1]) / (1.0 - self.abar[1:]) * beta def weight(self, which="tilde"): """ELBO 第 t 项里 ||eps - eps_theta||^2 前面的系数(长度 T,下标 0 对应 t=1)。 which="tilde" -> sigma^2 = tilde_beta_t (fixed_small) which="large" -> sigma^2 = beta_t (fixed_large) t=1 时 tilde_beta_1 = 0,权重分母为零(分子非零)。DDPM 把这一项交给 L_0(离散解码器) 单独处理;这里为了让训练能跑,退化地用 fixed_large 的权重顶上。 """ s2 = self.bt[1:] if which == "tilde" else self.beta with np.errstate(divide="ignore", invalid="ignore"): w = self.beta ** 2 / (2.0 * s2 * self.alpha * (1.0 - self.abar[1:])) if which == "tilde": w[0] = self.beta[0] / (2.0 * self.alpha[0] * (1.0 - self.abar[1])) return w def weight_simplified(self, which="tilde"): """上面那条式子化简之后的样子,用来核对代数没推错。""" denom = 1.0 - (self.abar[:-1] if which == "tilde" else self.abar[1:]) with np.errstate(divide="ignore", invalid="ignore"): return self.beta / (2.0 * self.alpha * denom) def q_sample(x0, t_idx, abar, rng): """x_t = sqrt(abar_t) x_0 + sqrt(1 - abar_t) eps,一步到位。""" a = abar[t_idx] if np.isscalar(a): a = np.full(len(x0), a) eps = rng.standard_normal(x0.shape) return np.sqrt(a)[:, None] * x0 + np.sqrt(1.0 - a)[:, None] * eps, eps # ══════════════════════════════════════════════════════════════════ # 1. 数据:一个二维高斯(NLL 能算到机器精度) # ══════════════════════════════════════════════════════════════════ MU0 = np.array([1.0, -0.5]) S0 = np.array([[0.60, 0.25], [0.25, 0.35]]) def sample_gauss(n, rng): L = np.linalg.cholesky(S0) return MU0 + rng.standard_normal((n, D)) @ L.T def entropy_gauss(S): """高斯熵(nats)。""" _, ld = np.linalg.slogdet(S) return 0.5 * (D * np.log(2 * np.pi * np.e) + ld) def kl_gauss(m0, S0_, m1, S1_): """KL( N(m0, S0_) || N(m1, S1_) ),闭式解。""" L = np.linalg.cholesky(S1_) y = np.linalg.solve(L, m1 - m0) Sinv_S0 = np.linalg.solve(S1_, S0_) _, ld1 = np.linalg.slogdet(S1_) _, ld0 = np.linalg.slogdet(S0_) return 0.5 * (np.trace(Sinv_S0) + y @ y - D + ld1 - ld0) def expected_nll(m, S): """E_{x ~ p_data}[-log p_theta(x)] 的**精确值**。 = 数据分布的熵 + KL(p_data || p_theta),两个高斯之间全是闭式解, 没有蒙特卡洛噪声——用固定测试集估 NLL 时,2 万样本的波动有 ±0.007 nats, 比我们要比的 0.002~0.004 nats 还大,所以这里必须用闭式解。 """ return entropy_gauss(S0) + kl_gauss(MU0, S0, m, S) def gauss_nll(X, m, S): """多元高斯负对数似然(nats)。协方差病态时加抖动兜底。""" jitter = 0.0 for _ in range(8): try: L = np.linalg.cholesky(S + jitter * np.eye(D)) y = np.linalg.solve(L, (X - m).T) return 0.5 * (np.sum(y * y, axis=0) + 2 * np.sum(np.log(np.diag(L))) + D * np.log(2 * np.pi)) except np.linalg.LinAlgError: jitter = 1e-12 if jitter == 0.0 else jitter * 100.0 raise RuntimeError("协方差矩阵无法正定化,说明反向链已经数值发散") # ══════════════════════════════════════════════════════════════════ # 2. 模型:eps_theta(x, t) = B(psi_t) x + c(psi_t),参数跨 t 共享 # psi 是固定的时间基(Fourier),只有 B、c 的系数是可学的 # ══════════════════════════════════════════════════════════════════ # 时间基故意取得很小:让「容量」成为真正的约束,权重才会决定容量往哪搬。 # 指数里出现 0.5,是因为最优的 B_t 与 sqrt(1 - abar_t) 成正比, # 在 t 很小时它按 sqrt(t) 走 —— 用纯多项式去拟合会在开头剧烈震荡。 BASIS = { "loose": [0.0, 0.5, 1.0, 1.5, 2.0], # 5 项:容量基本够用 "tight": [0.0, 0.5, 1.0], # 3 项:容量真的成了瓶颈 } BASIS_EXP = BASIS["tight"] # 默认用容量吃紧那一档,机制看得最清楚 K = len(BASIS_EXP) def set_basis(name): global BASIS_EXP, K BASIS_EXP = BASIS[name] K = len(BASIS_EXP) return K def psi_basis(t_arr): """t (1..T) -> (n, K) 的固定时间基,u = t / T in (0, 1]。""" u = np.asarray(t_arr, dtype=float) / T return np.stack([u ** e for e in BASIS_EXP], axis=-1) def feats(x, psi): """拼特征:[x ⊗ psi, psi],最后一维长度 D*K + K。""" px = x[:, :, None] * psi[:, None, :] # (n, D, K) return np.concatenate([px.reshape(len(x), -1), psi], axis=1) def fit_weighted_ridge(Phi, E, w, lam): """最小化 sum_i w_i ||W phi_i - eps_i||^2 + lam ||W||^2。""" P = Phi.shape[1] A = np.einsum("n,np,nq->pq", w, Phi, Phi) + lam * np.eye(P) Bmat = np.einsum("n,np,nd->pd", w, Phi, E) return np.linalg.solve(A, Bmat) def model_Bc(W, psi_t): """把共享权重在某个时刻 t 上展开成 eps = B x + c。""" Wx = W[: D * K, :].reshape(D, K, D) # Wx[i, k, j] Bmat = np.einsum("ikj,k->ji", Wx, psi_t) cvec = psi_t @ W[D * K:, :] return Bmat, cvec def optimal_Bc(sc, t): """高斯数据下 eps 的最优预测器(后验均值),用来做参考与验算。""" a = sc.abar[t] C = a * S0 + (1.0 - a) * np.eye(D) Bmat = np.sqrt(1.0 - a) * np.linalg.inv(C) cvec = -np.sqrt(a) * (Bmat @ MU0) return Bmat, cvec def mse_analytic(Bmat, cvec, sc, t): """E||B x_t + c - eps||^2 的闭式解(对 D 维取了平均)。 (x_t, eps) 是联合高斯:Var(x_t) = C_t = abar_t S0 + (1-abar_t) I, Cov(x_t, eps) = sqrt(1-abar_t) I,E[x_t] = sqrt(abar_t) mu0,E[eps] = 0。 展开 ||B x_t + c - eps||^2 的期望即得下式,不用蒙特卡洛,没有抽样噪声。 """ a = sc.abar[t] C = a * S0 + (1.0 - a) * np.eye(D) m = np.sqrt(a) * MU0 quad = np.trace(Bmat @ C @ Bmat.T) # Var(B x_t) bias = np.sum((Bmat @ m + cvec) ** 2) # 均值没对上的部分 cross = 2.0 * np.sqrt(1.0 - a) * np.trace(Bmat) # -2 Cov(B x_t, eps) return float((quad + bias - cross + D) / D) def mse_floor(sc, t): """贝叶斯最优的 MSE 地板:Var(eps | x_t) 的迹除以 D。""" a = sc.abar[t] C = a * S0 + (1.0 - a) * np.eye(D) return float((D - (1.0 - a) * np.trace(np.linalg.inv(C))) / D) def learnability(sc, t): """这一时刻的 eps 里,有多大比例是能从 x_t 里看出来的。 R^2 = 1 - 地板 / Var(eps),Var(eps) 的每维是 1。 """ return 1.0 - mse_floor(sc, t) # ══════════════════════════════════════════════════════════════════ # 3. 精确 NLL:把高斯密度沿反向链往前传 # ══════════════════════════════════════════════════════════════════ def _s2_mat(sigma2, t): """每步注入的噪声协方差:既接受标量(各向同性),也接受 (T+1, D, D)。""" if sigma2.ndim == 1: return sigma2[t] * np.eye(D) return sigma2[t] def chain_moments_prop(sc, Bcfun, sigma2): """写法一:密度传播。x_{t-1} = M_t x_t + v_t + noise。""" m = np.zeros(D) S = np.eye(D) sig = np.atleast_1d(sigma2) for t in range(T, 0, -1): Bmat, cvec = Bcfun(t) kk = sc.beta[t - 1] / np.sqrt(1.0 - sc.abar[t]) M = (np.eye(D) - kk * Bmat) / np.sqrt(sc.alpha[t - 1]) v = -kk * cvec / np.sqrt(sc.alpha[t - 1]) m = M @ m + v S = M @ S @ M.T + _s2_mat(sig, t) return m, S def chain_moments_comp(sc, Bcfun, sigma2): """写法二:把整条链复合成一个仿射映射,再累加各步噪声。独立实现,用于交叉验证。""" m = np.zeros(D) S = np.zeros((D, D)) Pprev = np.eye(D) for t in range(1, T + 1): Bmat, cvec = Bcfun(t) kk = sc.beta[t - 1] / np.sqrt(1.0 - sc.abar[t]) M = (np.eye(D) - kk * Bmat) / np.sqrt(sc.alpha[t - 1]) v = -kk * cvec / np.sqrt(sc.alpha[t - 1]) m = m + Pprev @ v S = S + Pprev @ _s2_mat(sigma2, t) @ Pprev.T Pprev = Pprev @ M S = S + Pprev @ Pprev.T # x_T ~ N(0, I) 的那一坨 return m, S def sigma2_exact_posterior(sc): """真实后验 q(x_{t-1}|x_t) 的协方差(高斯数据下可算)。 = tilde_beta_t I + Var(mu_t | x_t) = tilde_beta_t I + (beta_t^2 / (alpha_t (1 - abar_t))) Var(eps | x_t) 第二项就是 DDPM 反向链扔掉的那部分:真实后验比 beta_tilde 更胖。 """ out = np.zeros((T + 1, D, D)) for t in range(1, T + 1): a = sc.abar[t] C = a * S0 + (1.0 - a) * np.eye(D) var_eps = np.eye(D) - (1.0 - a) * np.linalg.inv(C) kk2 = sc.beta[t - 1] ** 2 / (sc.alpha[t - 1] * (1.0 - a)) out[t] = sc.bt[t] * np.eye(D) + kk2 * var_eps return out def sigma2_from_choice(sc, which="tilde", lam_mix=0.0): """反向链每步注入的方差。lam_mix 在 tilde_beta 与 beta 之间插值。""" if which == "tilde": base = sc.bt.copy() other = np.concatenate([[0.0], sc.beta]) else: base = np.concatenate([[0.0], sc.beta]) other = sc.bt.copy() return (1.0 - lam_mix) * base + lam_mix * other # ══════════════════════════════════════════════════════════════════ # 4. 主流程 # ══════════════════════════════════════════════════════════════════ def section_A(sc): print("=" * 74) print("A. ELBO 每步权重:解析式 vs 两个高斯的 KL") print("=" * 74) w_tilde = sc.weight("tilde") w_large = sc.weight("large") ws_tilde = sc.weight_simplified("tilde") ws_large = sc.weight_simplified("large") ok = np.isfinite(ws_tilde) & (ws_tilde > 0) rel = np.max(np.abs(w_tilde[ok] - ws_tilde[ok]) / ws_tilde[ok]) print(f" 化简式与原始式的最大相对差(tilde):{rel:.3e}") rel2 = np.max(np.abs(w_large - ws_large) / ws_large) print(f" 化简式与原始式的最大相对差(large):{rel2:.3e}") print(f" t=1 的 tilde_beta_1 = {sc.bt[1]:.3e}(后验方差为 0 => 权重发散," f"这就是 DDPM 要把 L_0 单独拿出来的原因)") print() print(" t abar_t beta_t w_t(sigma^2=bt) w_t(sigma^2=beta) 相对等权") for t in [2, 10, 50, 100, 250, 500, 750, 900, 1000]: print(f" {t:<6d}{sc.abar[t]:<12.3e}{sc.beta[t-1]:<10.5f}" f"{w_tilde[t-1]:<18.4e}{w_large[t-1]:<20.4e}" f"{w_tilde[t-1] / w_tilde[499]:<10.2f}") print() span = w_tilde[1] / w_tilde.min() print(f" sigma^2=tilde_beta 时,权重最大值(t=2)是最小值(t={int(np.argmin(w_tilde))+1})的 " f"{span:.1f} 倍") print(f" 等权 L_simple 相当于把这条曲线整体压平到 1") print() # 数值核对:KL( q(x_{t-1}|x_t,x_0) || p_theta(x_{t-1}|x_t) ) == w_t * ||delta_eps||^2 rng = np.random.default_rng(RNG_SEED) print(" —— KL 数值核对(随机取 x_0, eps,令 eps_theta = eps + delta)——") print(" t KL(数值) w_t*||delta||^2 相对差") for t in [2, 100, 500, 1000]: x0 = sample_gauss(1, rng)[0] a = sc.abar[t] eps = rng.standard_normal(D) xt = np.sqrt(a) * x0 + np.sqrt(1 - a) * eps delta = rng.standard_normal(D) * 0.3 eps_hat = eps + delta kk = sc.beta[t - 1] / np.sqrt(1 - a) mu_tilde = (xt - kk * eps) / np.sqrt(sc.alpha[t - 1]) mu_theta = (xt - kk * eps_hat) / np.sqrt(sc.alpha[t - 1]) s2 = sc.bt[t] kl = np.sum((mu_tilde - mu_theta) ** 2) / (2 * s2) pred = w_tilde[t - 1] * np.sum(delta ** 2) print(f" {t:<7d}{kl:<18.6e}{pred:<20.6e}{abs(kl-pred)/kl:<12.2e}") print() return w_tilde, w_large def fit_three(sc, w_tilde, w_large, lam=1e-6): """三种权重各解一遍带权最小二乘。不打印任何东西,供 section_B / section_F 共用。""" rng = np.random.default_rng(RNG_SEED + 1) M = 120_000 t_idx = rng.integers(1, T + 1, size=M) x0 = sample_gauss(M, rng) a = sc.abar[t_idx] eps = rng.standard_normal((M, D)) xt = np.sqrt(a)[:, None] * x0 + np.sqrt(1 - a)[:, None] * eps psi = psi_basis(t_idx) Phi = feats(xt, psi) # 三种权重都归一化到均值 1(整体缩放不改变无正则时的解, # 归一化只是为了让正则强度在三种目标下可比) schemes = { "uniform": np.ones(T), "elbo_tilde": w_tilde / w_tilde.mean(), "elbo_large": w_large / w_large.mean(), } lam = 1e-6 W = {} for name, w_full in schemes.items(): w = w_full[t_idx - 1] W[name] = fit_weighted_ridge(Phi, eps, w, lam) return W, schemes def section_B(sc, w_tilde, w_large): print("=" * 74) print("B. 带权最小二乘训练:三种目标,共享参数,容量有限") print("=" * 74) W, schemes = fit_three(sc, w_tilde, w_large) print(f" 训练集:M=120000 条 (x_0, t, eps),时间基 {K} 项,特征维度 {K * (D + 1)}," f"可学参数 {K * (D + 1) * D} 个") print(f" 三种权重都归一化到均值 1(整体缩放不改变解,归一化只为让正则强度可比)") print(f" 权重动态范围:elbo_tilde 的 max/min = {w_tilde.max()/w_tilde.min():.1f}") print() return W, schemes def model_moments(sc, Wm, which="tilde"): """把训练好的共享权重展开成整条反向链→后验的均值与协方差。""" s2 = sigma2_from_choice(sc, which) return chain_moments_prop( sc, lambda t: model_Bc(Wm, psi_basis(np.array([t]))[0]), s2) def excess_bands(sc, W): """低噪声档 / 高噪声档的超额误差均值,以及 elbo 相对 uniform 的比值。""" lo = np.arange(1, 51) hi = np.arange(400, 1001) out = {} for band, idx in [("lo", lo), ("hi", hi)]: e = {} for name in ["uniform", "elbo_tilde"]: e[name] = float(np.mean([ mse_analytic(*model_Bc(W[name], psi_basis(np.array([t]))[0]), sc, t) - mse_floor(sc, t) for t in idx])) out[band] = (e["uniform"], e["elbo_tilde"], e["elbo_tilde"] / e["uniform"]) return out def section_C(sc, W, schemes): print("=" * 74) print("C. 仿射模型下的精确 NLL(密度前向传播)") print("=" * 74) rng = np.random.default_rng(RNG_SEED + 2) Xtest = sample_gauss(400_000, rng) nll_true = expected_nll(MU0, S0) print(f" 参考:真分布 N(mu_0, S_0) 的熵 = {nll_true:.6f} nats" f"({nll_true / np.log(2) / D:.4f} bits/dim)") print(f" 蒙特卡洛复核(40 万样本):{gauss_nll(Xtest, MU0, S0).mean():.6f} nats" f",与闭式解差 {abs(gauss_nll(Xtest, MU0, S0).mean()-nll_true):.2e}") # 交叉验证两种写法 m1, S1 = chain_moments_prop(sc, lambda t: optimal_Bc(sc, t), sigma2_from_choice(sc, "tilde")) m2, S2 = chain_moments_comp(sc, lambda t: optimal_Bc(sc, t), sigma2_from_choice(sc, "tilde")) print(f" 两种写法的差异:均值 {np.max(np.abs(m1-m2)):.2e},协方差 {np.max(np.abs(S1-S2)):.2e}") # 蒙特卡洛复核:真跑 20 万条链,比对经验均值/协方差 rng3 = np.random.default_rng(RNG_SEED + 3) xs = rng3.standard_normal((200_000, D)) for t in range(T, 0, -1): Bmat, cvec = optimal_Bc(sc, t) kk = sc.beta[t - 1] / np.sqrt(1.0 - sc.abar[t]) mu = (xs - kk * (xs @ Bmat.T + cvec)) / np.sqrt(sc.alpha[t - 1]) xs = mu + np.sqrt(sc.bt[t]) * rng3.standard_normal((200_000, D)) print(f" 蒙特卡洛复核:均值差 {np.max(np.abs(xs.mean(0)-m1)):.2e}," f"协方差差 {np.max(np.abs(np.cov(xs.T)-S1)):.2e}") nll_opt = expected_nll(m1, S1) print(f" 「最优 eps 预测器 + sigma^2=tilde_beta」的期望 NLL = {nll_opt:.6f} nats," f"比真分布差 {nll_opt - nll_true:+.6f} nats") print(f" 蒙特卡洛复核(同上 40 万样本):{gauss_nll(Xtest, m1, S1).mean():.6f} nats," f"差 {abs(gauss_nll(Xtest, m1, S1).mean()-nll_opt):.2e}") print(f" 生成分布的协方差 diag = {np.diag(S1)},真值 diag = {np.diag(S0)} " f"=> 比值 {np.diag(S1)/np.diag(S0)}") # 把真实后验方差(而不是 tilde_beta)灌回反向链:应该精确还原数据分布 s2_exact = sigma2_exact_posterior(sc) m3, S3 = chain_moments_prop(sc, lambda t: optimal_Bc(sc, t), s2_exact) nll_exact = expected_nll(m3, S3) print(f" 「最优 eps 预测器 + 真实后验方差」的期望 NLL = {nll_exact:.6f} nats," f"比真分布差 {nll_exact - nll_true:+.2e} nats") print(f" 此时生成协方差 diag / 真值 = {np.diag(S3)/np.diag(S0)}," f"均值差 {np.max(np.abs(m3 - MU0)):.2e}") print(" => reverse 链的均值用最优 eps 预测器、方差用真实后验方差," "就能精确还原数据分布;") print(" NLL 上剩下的那 0.0001 nats 完全是「把方差钉死成 tilde_beta」造成的。") print() print(" 三种训练目标 × 两种采样方差的期望 NLL(nats,越小越好):") print(" 训练目标 sigma^2=tilde_beta sigma^2=beta 差 生成std/真std") rows = {} for name in ["uniform", "elbo_tilde", "elbo_large"]: Wm = W[name] out = {} for which in ["tilde", "large"]: m, S = model_moments(sc, Wm, which) out[which] = expected_nll(m, S) out[which + "_std"] = float(np.mean(np.sqrt(np.diag(S)) / np.sqrt(np.diag(S0)))) rows[name] = out print(f" {name:<16}{out['tilde']:<22.6f}{out['large']:<18.6f}" f"{out['large']-out['tilde']:<+12.6f}{out['tilde_std']:.4f}") print() return rows, nll_true, nll_opt def section_D(sc, W): print("=" * 74) print("D. 每步误差剖面:权重把「超额误差」搬到了哪") print("=" * 74) grid = list(range(1, 1001)) floor = np.array([mse_floor(sc, t) for t in grid]) learn = np.array([learnability(sc, t) for t in grid]) prof = {} for name, Wm in W.items(): mse = [] for t in grid: Bmat, cvec = model_Bc(Wm, psi_basis(np.array([t]))[0]) mse.append(mse_analytic(Bmat, cvec, sc, t)) prof[name] = np.array(mse) print(" 闭式解算的,没有蒙特卡洛噪声。『超额』= MSE − 贝叶斯地板。") print(" t 可学占比R^2 地板 uniform超额 elbo_tilde超额 elbo_large超额") for t in [1, 5, 20, 50, 100, 200, 400, 600, 800, 950, 1000]: i = t - 1 print(f" {t:<6d}{learn[i]:<12.4f}{floor[i]:<9.4f}" f"{prof['uniform'][i]-floor[i]:<14.3e}{prof['elbo_tilde'][i]-floor[i]:<16.3e}" f"{prof['elbo_large'][i]-floor[i]:.3e}") print() lo = slice(0, 50) # t = 1..50 hi = slice(399, 1000) # t = 400..1000 eu_u, eu_e = prof["uniform"] - floor, prof["elbo_tilde"] - floor print(f" 低噪声档 t<=50 :uniform 超额均值 {eu_u[lo].mean():.3e}," f"elbo_tilde {eu_e[lo].mean():.3e}(比值 {eu_e[lo].mean()/eu_u[lo].mean():.2f})") print(f" 高噪声档 t>=400 :uniform 超额均值 {eu_u[hi].mean():.3e}," f"elbo_tilde {eu_e[hi].mean():.3e}(比值 {eu_e[hi].mean()/eu_u[hi].mean():.2f})") print() print(f" 可学占比 R^2:t=1 时 {learn[0]:.2e},t=500 时 {learn[499]:.4f}," f"t=1000 时 {learn[999]:.4f}") print() return grid, prof, floor, learn def section_E(sc, W): print("=" * 74) print("E. 采样方差该取 tilde_beta 还是 beta:在两者之间插值扫一遍") print("=" * 74) nll_true = expected_nll(MU0, S0) print(" lam 是插值系数:sigma^2 = (1-lam)*tilde_beta + lam*beta") print(" lam NLL(uniform训练) NLL(elbo_tilde训练) 生成 std / 真 std") out = {"lam": [], "nll_uniform": [], "nll_elbo": [], "std_ratio": []} for lam in [0.0, 0.25, 0.5, 0.75, 1.0]: s2 = sigma2_from_choice(sc, "tilde", lam_mix=lam) row, sr = [], 0.0 for name in ["uniform", "elbo_tilde"]: m, S = chain_moments_prop( sc, lambda t: model_Bc(W[name], psi_basis(np.array([t]))[0]), s2) row.append(expected_nll(m, S)) if name == "uniform": sr = float(np.mean(np.sqrt(np.diag(S)) / np.sqrt(np.diag(S0)))) out["lam"].append(lam) out["nll_uniform"].append(row[0]) out["nll_elbo"].append(row[1]) out["std_ratio"].append(sr) print(f" {lam:<7.2f}{row[0]:<19.6f}{row[1]:<22.6f}{sr:.4f}") best = out["lam"][int(np.argmin(np.abs(np.array(out["std_ratio"]) - 1.0)))] print(f" 生成分布的胖瘦刚好对上真分布的 lam ≈ {best:.2f}(这一列是最敏感的指标," f"NLL 对 lam 几乎不动)") print(f" 参考:真分布熵 = {nll_true:.6f}") print() return out def section_F(sc, w_tilde, w_large): """容量充足 vs 容量吃紧:『该不该扔掉系数』的答案会不会反过来。""" print("=" * 74) print("F. 同样两个目标,换一档模型容量:结论会反过来") print("=" * 74) print(" 两档时间基都只动特征个数,训练数据、种子、正则强度完全一样。") print() print(" 容量档 时间基项数 可学参数 NLL(等权) NLL(真ELBO) 谁更好 低噪声档超额比 高噪声档超额比") out = {} for regime in ["loose", "tight"]: k = set_basis(regime) Wf, _ = fit_three(sc, w_tilde, w_large) n_u = expected_nll(*model_moments(sc, Wf["uniform"])) n_e = expected_nll(*model_moments(sc, Wf["elbo_tilde"])) bands = excess_bands(sc, Wf) better = "等权" if n_u < n_e else "真ELBO" n_par = k * (D + 1) * D out[regime] = dict(K=k, nll_uniform=n_u, nll_elbo=n_e, bands=bands, W=Wf) print(f" {regime:<8}{k:<12}{n_par:<11}{n_u:<12.6f}{n_e:<14.6f}" f"{better:<9}{bands['lo'][2]:<16.2f}{bands['hi'][2]:.2f}") set_basis("tight") print() print(" 列『超额比』= 真 ELBO 权重的超额误差 / 等权的超额误差,小于 1 表示更好。") print(" => 容量够用时真 ELBO 权重略胜,容量真的吃紧时它反而输给等权。") print() return out def main(): sc = Sched(linear_beta(T)) print(f"调度:线性 T={T},beta 从 {sc.beta[0]:.1e} 到 {sc.beta[-1]:.3f}," f"abar_T = {sc.abar[T]:.4e}") print() w_tilde, w_large = section_A(sc) W, schemes = section_B(sc, w_tilde, w_large) print() rows, nll_true, nll_opt = section_C(sc, W, schemes) print() grid, prof, floor, learn = section_D(sc, W) print() mix = section_E(sc, W) print() regimes = section_F(sc, w_tilde, w_large) return dict(sched=sc, w_tilde=w_tilde, w_large=w_large, W=W, nll_rows=rows, nll_true=nll_true, nll_opt=nll_opt, grid=grid, prof=prof, floor=floor, learn=learn, mix=mix, regimes=regimes) if __name__ == "__main__": main() ddpm_train.py # -*- coding: utf-8 -*- """照着 DDPM 原文 Algorithm 1 / Algorithm 2 写的最小实现(纯 numpy,无 torch)。 为什么还要写一遍神经网络版:ddpm_lab.py 用带权最小二乘的闭式解把「优化算法」 这个变量消掉了,代价是模型只能仿射。这里补上手写反向传播的两层 MLP, 在**多模态**的目标分布上跑完整的一千步采样,看训练目标的权重到底怎么影响 最后生成出来的东西。 A. 梯度核对:手写反向传播 vs 有限差分 B. Algorithm 1:按三种加权(等权 / ELBO / ELBO-large)训练三个网络 C. Algorithm 2:一千步采样,三种指标打分 D. 反向链方差 sigma^2 = tilde_beta 还是 beta 运行: /usr/local/bin/python3 ddpm_train.py # 默认 20000 步 /usr/local/bin/python3 ddpm_train.py --steps 40000 依赖: numpy """ import argparse import os import numpy as np RNG_SEED = 20260927 T = 1000 D = 2 # ══════════════════════════════════════════════════════════════════ # 0. 调度(与 ddpm_lab.py 同一套,下标约定:t 取 1..T) # ══════════════════════════════════════════════════════════════════ def linear_beta(T=1000, b1=1e-4, bT=0.02): return np.linspace(b1, bT, T) class Sched: def __init__(self, beta): self.beta = beta self.alpha = 1.0 - beta self.abar = np.concatenate([[1.0], np.cumprod(self.alpha)]) self.bt = np.empty(len(beta) + 1) self.bt[0] = 0.0 self.bt[1:] = (1.0 - self.abar[:-1]) / (1.0 - self.abar[1:]) * beta def weight(self, which="tilde"): """ELBO 第 t 项里 ||eps - eps_theta||^2 的系数。""" s2 = self.bt[1:] if which == "tilde" else self.beta with np.errstate(divide="ignore", invalid="ignore"): w = self.beta ** 2 / (2.0 * s2 * self.alpha * (1.0 - self.abar[1:])) if which == "tilde": # t=1 的 0/0 退化,用 fixed_large 顶上 w[0] = self.beta[0] / (2.0 * self.alpha[0] * (1.0 - self.abar[1])) return w # ══════════════════════════════════════════════════════════════════ # 1. 目标分布:7 个分量的二维高斯混合(环状 + 一个中心) # ══════════════════════════════════════════════════════════════════ RING_R = 1.15 N_RING = 6 CENTERS = np.array( [[RING_R * np.cos(2 * np.pi * k / N_RING), RING_R * np.sin(2 * np.pi * k / N_RING)] for k in range(N_RING)] + [[0.0, 0.0]] ) SCALES = np.array([0.17, 0.20, 0.16, 0.22, 0.18, 0.19, 0.28]) WEIGHTS = SCALES ** 2 # 分量越大越容易被采到,制造不均匀 WEIGHTS = WEIGHTS / WEIGHTS.sum() def sample_mix(n, rng): kk = rng.choice(len(WEIGHTS), size=n, p=WEIGHTS) return CENTERS[kk] + SCALES[kk][:, None] * rng.standard_normal((n, D)) def logp_mix(X): """混合分布在 X 处的对数密度(各分量都是各向同性高斯)。""" d2 = ((X[:, None, :] - CENTERS[None, :, :]) ** 2).sum(-1) # (n, K) comp = np.log(WEIGHTS)[None, :] - 0.5 * d2 / SCALES[None, :] ** 2 \ - D * np.log(SCALES[None, :]) - 0.5 * D * np.log(2 * np.pi) m = comp.max(axis=1, keepdims=True) return (m[:, 0] + np.log(np.exp(comp - m).sum(axis=1))) def nearest_center_dist(X): d2 = ((X[:, None, :] - CENTERS[None, :, :]) ** 2).sum(-1) return np.sqrt(d2.min(axis=1)) # ══════════════════════════════════════════════════════════════════ # 2. 时间嵌入 + 两层 MLP + 手写 Adam # ══════════════════════════════════════════════════════════════════ TIME_FREQ = (1, 2, 4, 8) # 8 维时间基 def time_embed(t_arr): u = np.asarray(t_arr, dtype=float) / T cols = [] for w in TIME_FREQ: cols.append(np.cos(w * np.pi * u)) cols.append(np.sin(w * np.pi * u)) return np.stack(cols, axis=-1) N_IN = D + 2 * len(TIME_FREQ) H = 64 def init_params(rng): def he(fan_in, fan_out): return rng.standard_normal((fan_in, fan_out)) * np.sqrt(2.0 / fan_in) return { "W1": he(N_IN, H), "b1": np.zeros(H), "W2": he(H, H), "b2": np.zeros(H), "W3": he(H, D), "b3": np.zeros(D), } def forward(p, Z): h1 = np.maximum(Z @ p["W1"] + p["b1"], 0.0) h2 = np.maximum(h1 @ p["W2"] + p["b2"], 0.0) return h1, h2, h2 @ p["W3"] + p["b3"] def backward(p, Z, h1, h2, out, eps, w): """d/d(theta) of mean_i w_i ||out_i - eps_i||^2。""" n = len(Z) dout = 2.0 * (w[:, None] * (out - eps)) / n g = {} g["W3"] = h2.T @ dout g["b3"] = dout.sum(0) dh2 = dout @ p["W3"].T dh2[h2 <= 0] = 0.0 g["W2"] = h1.T @ dh2 g["b2"] = dh2.sum(0) dh1 = dh2 @ p["W2"].T dh1[h1 <= 0] = 0.0 g["W1"] = Z.T @ dh1 g["b1"] = dh1.sum(0) return g class Adam: def __init__(self, p, lr=1e-3, b1=0.9, b2=0.999, eps=1e-8): self.m = {k: np.zeros_like(v) for k, v in p.items()} self.v = {k: np.zeros_like(v) for k, v in p.items()} self.lr, self.b1, self.b2, self.eps = lr, b1, b2, eps self.i = 0 def step(self, p, g): self.i += 1 for k in p: self.m[k] = self.b1 * self.m[k] + (1 - self.b1) * g[k] self.v[k] = self.b2 * self.v[k] + (1 - self.b2) * g[k] ** 2 mh = self.m[k] / (1 - self.b1 ** self.i) vh = self.v[k] / (1 - self.b2 ** self.i) p[k] -= self.lr * mh / (np.sqrt(vh) + self.eps) def predict_eps(p, x, t_idx): Z = np.concatenate([x, time_embed(np.full(len(x), t_idx))], axis=1) _, _, out = forward(p, Z) return out # ══════════════════════════════════════════════════════════════════ # 3. Algorithm 1(训练)与 Algorithm 2(采样) # ══════════════════════════════════════════════════════════════════ def train(sc, w_full, steps, batch, rng, lr=1e-3): """DDPM 原文 Algorithm 1,一行不差地照抄。""" p = init_params(rng) opt = Adam(p, lr=lr) for _ in range(steps): # 1: t ~ Uniform({1, ..., T}) 2: x_0 ~ q(x_0) 3: eps ~ N(0, I) t_idx = rng.integers(1, T + 1, size=batch) x0 = sample_mix(batch, rng) eps = rng.standard_normal((batch, D)) # 4: 一步加噪(闭式解,不需要真的走 t 步) a = sc.abar[t_idx] xt = np.sqrt(a)[:, None] * x0 + np.sqrt(1.0 - a)[:, None] * eps # 5: 梯度下降一步 Z = np.concatenate([xt, time_embed(t_idx)], axis=1) h1, h2, out = forward(p, Z) opt.step(p, backward(p, Z, h1, h2, out, eps, w_full[t_idx - 1])) return p def sample_chain(sc, p, n, which="tilde", rng=None): """DDPM 原文 Algorithm 2。which 决定 sigma^2 取 tilde_beta 还是 beta。""" x = rng.standard_normal((n, D)) # x_T ~ N(0, I) for t in range(T, 0, -1): eps_hat = predict_eps(p, x, t) a, a_prev = sc.abar[t], sc.abar[t - 1] alpha_t, beta_t = sc.alpha[t - 1], sc.beta[t - 1] # mu_tilde = (x_t - beta_t/sqrt(1-abar_t) * eps_hat) / sqrt(alpha_t) mu = (x - beta_t / np.sqrt(1.0 - a) * eps_hat) / np.sqrt(alpha_t) if which == "tilde": sigma = np.sqrt(sc.bt[t]) else: sigma = np.sqrt(beta_t) z = rng.standard_normal((n, D)) x = mu + (sigma * z if t > 1 else 0.0) # t == 1 时不加噪声 return x # ══════════════════════════════════════════════════════════════════ # 4. 主流程 # ══════════════════════════════════════════════════════════════════ def section_A(): print("=" * 74) print("A. 手写反向传播 vs 有限差分") print("=" * 74) rng = np.random.default_rng(RNG_SEED) sc = Sched(linear_beta(T)) p = init_params(rng) n = 32 t_idx = rng.integers(1, T + 1, size=n) x0 = sample_mix(n, rng) eps = rng.standard_normal((n, D)) a = sc.abar[t_idx] xt = np.sqrt(a)[:, None] * x0 + np.sqrt(1 - a)[:, None] * eps Z = np.concatenate([xt, time_embed(t_idx)], axis=1) w = np.ones(n) h1, h2, out = forward(p, Z) g = backward(p, Z, h1, h2, out, eps, w) def loss(): _, _, o = forward(p, Z) return float(np.mean(np.sum((o - eps) ** 2, axis=1))) print(" 参数 解析梯度 有限差分 相对差") worst = 0.0 for key in ["W1", "b1", "W2", "b2", "W3", "b3"]: idx = tuple(0 for _ in p[key].shape) h = 1e-6 orig = p[key][idx] p[key][idx] = orig + h lp = loss() p[key][idx] = orig - h lm = loss() p[key][idx] = orig num = (lp - lm) / (2 * h) ana = g[key][idx] rel = abs(num - ana) / max(abs(num), 1e-12) worst = max(worst, rel) print(f" {key:<8}{ana:<16.8e}{num:<16.8e}{rel:.2e}") print(f" => 最大相对差 {worst:.2e}(有限差分自己的精度极限在 1e-6 量级)") print() return sc def section_B(sc, steps): print("=" * 74) print("B. Algorithm 1:三种加权各训一个网络") print("=" * 74) w_tilde, w_large = sc.weight("tilde"), sc.weight("large") schemes = { "uniform": np.ones(T), "elbo_tilde": w_tilde / w_tilde.mean(), } print(f" 步数 {steps},batch 256,两层 MLP({N_IN} -> {H} -> {H} -> {D}),Adam lr=1e-3") print(f" elbo 权重的动态范围:max/min = {w_tilde.max()/w_tilde.min():.1f}") models = {} # 注意:不能用 hash(name)——Python 的字符串 hash 每次进程都变,结果会不可复现。 seed_of = {"uniform": 11, "elbo_tilde": 22, "elbo_large": 33} for name, w in schemes.items(): rng = np.random.default_rng(RNG_SEED + seed_of.get(name, 44)) models[name] = train(sc, w, steps, 256, rng) print(f" [{name}] 训练完成") print() return models def _score(X, base_logp, base_dist): """三个指标:log p0 相对真样本、到最近模式中心的距离比、模式覆盖数。""" dlogp = logp_mix(X).mean() - base_logp dratio = nearest_center_dist(X).mean() / base_dist assign = ((X[:, None, :] - CENTERS[None, :, :]) ** 2).sum(-1).argmin(1) frac = np.bincount(assign, minlength=len(CENTERS)) / len(X) return dlogp, dratio, int((frac > 0.02).sum()) def section_C(sc, models, n_gen, n_seed=3): print("=" * 74) print("C. Algorithm 2:一千步采样,三个指标打分") print("=" * 74) rng0 = np.random.default_rng(RNG_SEED + 77) Xtrue = sample_mix(20_000, rng0) base_logp = logp_mix(Xtrue).mean() base_dist = nearest_center_dist(Xtrue).mean() print(f" 基线(2 万真样本):log p0 均值 {base_logp:.4f}," f"到最近模式中心距离均值 {base_dist:.4f}") print(f" 每组用 {n_seed} 个不同的采样种子重复,报告均值 ± 标准差") print() print(" 训练目标 log p0 相对真样本 到最近中心距离/真样本 模式覆盖") res = {} for name, p in models.items(): dl, dr, cov = [], [], [] for s in range(n_seed): rng = np.random.default_rng(RNG_SEED + 88 + s) a, b, c = _score(sample_chain(sc, p, n_gen, "tilde", rng), base_logp, base_dist) dl.append(a) dr.append(b) cov.append(c) res[name] = dict(dlogp=float(np.mean(dl)), dlogp_std=float(np.std(dl, ddof=1)), dratio=float(np.mean(dr)), frac=cov[0]) print(f" {name:<14}{np.mean(dl):<+10.4f} ± {np.std(dl, ddof=1):<10.4f}" f"{np.mean(dr):<22.4f}{cov[0]}/7") print() return res, base_logp, base_dist def section_D(sc, models, n_gen, n_seed=3): print("=" * 74) print("D. 反向链方差:tilde_beta(fixed_small)还是 beta(fixed_large)") print("=" * 74) rng0 = np.random.default_rng(RNG_SEED + 77) Xtrue = sample_mix(20_000, rng0) base_logp = logp_mix(Xtrue).mean() base_dist = nearest_center_dist(Xtrue).mean() p = models["uniform"] print(" 用 uniform(L_simple)训出来的那个网络,只换每步注入的方差:") print(" 方差选择 log p0 相对真样本 到最近中心距离/真样本") out = {} for which in ["tilde", "large"]: dl, dr = [], [] for s in range(n_seed): rng = np.random.default_rng(RNG_SEED + 99 + s) a, b, _ = _score(sample_chain(sc, p, n_gen, which, rng), base_logp, base_dist) dl.append(a) dr.append(b) out[which] = dict(dlogp=float(np.mean(dl)), dlogp_std=float(np.std(dl, ddof=1)), dratio=float(np.mean(dr))) print(f" {which:<14}{np.mean(dl):<+10.4f} ± {np.std(dl, ddof=1):<10.4f}{np.mean(dr):.4f}") print() return out def mse_profile(sc, models, rng_seed=RNG_SEED + 123): """每个时刻的 eps 预测误差(蒙特卡洛,用来画图)。""" rng = np.random.default_rng(rng_seed) grid = np.array([1, 2, 5, 10, 20, 50, 100, 200, 400, 600, 800, 950, 1000]) n = 6000 prof = {} for name, p in models.items(): vals = [] for t in grid: x0 = sample_mix(n, rng) a = sc.abar[t] eps = rng.standard_normal((n, D)) xt = np.sqrt(a) * x0 + np.sqrt(1 - a) * eps vals.append(float(np.mean((predict_eps(p, xt, t) - eps) ** 2))) prof[name] = np.array(vals) return grid, prof def main(): ap = argparse.ArgumentParser() ap.add_argument("--steps", type=int, default=20000) ap.add_argument("--n-gen", type=int, default=4000) args = ap.parse_args() sc = section_A() models = section_B(sc, args.steps) res, base_logp, base_dist = section_C(sc, models, args.n_gen) dres = section_D(sc, models, args.n_gen) grid, prof = mse_profile(sc, models) print("=" * 74) print("E. 每步预测误差(MSE,真 eps 的每维方差是 1)") print("=" * 74) print(" t uniform elbo_tilde") for i, t in enumerate(grid): print(f" {t:<7d}{prof['uniform'][i]:<12.4f}{prof['elbo_tilde'][i]:.4f}") return dict(sched=sc, models=models, res=res, dres=dres, grid=grid, prof=prof) if __name__ == "__main__": main() make_figures.py # -*- coding: utf-8 -*- """画本文的四张图。数据源全部来自 ddpm_lab.py 的真实输出,不另造数。 weight_profile.png ELBO 每步权重 vs 这个时刻真正可学的信号占比 excess_error.png 两种训练目标把「超额误差」搬到了哪 variance_ledger.png 反向链每步注入的噪声:beta_tilde / beta / 真实后验 step_anatomy.png 反向一步的三个配料怎么随 t 变 运行: /usr/local/bin/python3 make_figures.py 依赖: numpy, matplotlib(字体 PingFang SC) 注意:matplotlib 的 mathtext 标签一律用 raw 字符串,且反斜杠后面只能跟字母 ——源码会被 sync-code 原样搬进文章附录,反斜杠后面跟非字母会被体检器判成转义污染。 """ import os import numpy as np import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import ddpm_lab as LAB 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) C_MAIN, C_ALT, C_GREEN, C_GRAY = "#2563eb", "#dc2626", "#059669", "#6b7280" T = LAB.T D = LAB.D # ─────────────────────────── 图 1:权重 vs 可学占比 ─────────────────────────── def fig_weight_profile(sc, w_tilde, w_large, learn): fig, ax = plt.subplots(figsize=(9.2, 5.4)) tt = np.arange(1, T + 1) ax.semilogy(tt, w_tilde, color=C_MAIN, lw=2.0, label=r"真 ELBO 权重 $w_t$($\sigma^2=\tilde\beta_t$)") ax.semilogy(tt, w_large, color=C_ALT, lw=1.6, ls="--", label=r"真 ELBO 权重 $w_t$($\sigma^2=\beta_t$)") ax.axhline(1.0, color=C_GRAY, lw=1.2, ls=":") ax.text(620, 1.15, "等权 $L_{\mathrm{simple}}$ 就压在这条 1 上", color=C_GRAY, fontsize=9) ax.set_xlabel(r"时间步 $t$") ax.set_ylabel(r"预测误差平方前面的系数 $w_t$") ax.set_ylim(3e-3, 3.0) ax.grid(alpha=0.25, which="both") ax.legend(fontsize=9, loc="lower left") ax2 = ax.twinx() ax2.plot(tt, learn, color=C_GREEN, lw=2.0) ax2.set_ylabel(r"这一时刻能学出来的噪声占比 $R^2$", color=C_GREEN) ax2.set_ylim(-0.05, 1.05) ax2.tick_params(axis="y", labelcolor=C_GREEN) ax2.set_title("(a) 权重最大地方,恰恰是最学不到东西的地方", fontsize=11) span = w_tilde.max() / w_tilde.min() ax.annotate(rf"权重跨 {span:.0f} 倍", xy=(2, w_tilde[0]), xytext=(70, 0.35), fontsize=9, color=C_MAIN, arrowprops=dict(arrowstyle="->", color=C_MAIN, lw=1.2)) ax.annotate(rf"$t=1$ 时 $R^2$只有 {learn[0]:.1e}", xy=(1, 0.02), xytext=(120, 0.012), fontsize=9, color=C_GREEN, arrowprops=dict(arrowstyle="->", color=C_GREEN, lw=1.2)) fig.tight_layout() fig.savefig(os.path.join(FIGDIR, "weight_profile.png"), dpi=130) plt.close(fig) # ──────────────────────── 图 2:超额误差被搬到哪 ──────────────────────── def fig_excess(sc, grid, prof, floor, learn): fig, axes = plt.subplots(1, 2, figsize=(13.2, 5.0)) tt = np.asarray(grid) ax = axes[0] ax.semilogy(tt, prof["uniform"] - floor, color=C_MAIN, lw=2.0, label=r"等权($L_{\mathrm{simple}}$)") ax.semilogy(tt, prof["elbo_tilde"] - floor, color=C_ALT, lw=2.0, label="真 ELBO 权重") ax.set_xlabel(r"时间步 $t$") ax.set_ylabel(r"超额误差 MSE $-$ 贝叶斯地板") ax.set_title("(a) 容量被搬走了:低噪声档变好,中高噪声档变差") ax.grid(alpha=0.25, which="both") ax.legend(fontsize=9) ax.set_xlim(0, 1000) ax = axes[1] ratio = (prof["elbo_tilde"] - floor) / (prof["uniform"] - floor) ax.semilogy(tt, ratio, color=C_GREEN, lw=2.0) ax.axhline(1.0, color=C_GRAY, lw=1.2, ls=":") ax.fill_between(tt, 1e-3, ratio, where=(ratio < 1.0), color=C_MAIN, alpha=0.13) ax.fill_between(tt, 1.0, ratio, where=(ratio > 1.0), color=C_ALT, alpha=0.13) ax.set_xlabel(r"时间步 $t$") ax.set_ylabel("真 ELBO 权重 / 等权 的超额误差之比") ax.set_title("(b) 同一条曲线取比值:1 以下变好,1 以上变差") ax.grid(alpha=0.25, which="both") ax.set_xlim(0, 1000) ax.set_ylim(1e-2, 1e2) ax.text(60, 0.022, "低噪声档:好 10 倍以上", color=C_MAIN, fontsize=9) ax.text(430, 4.5, "中高噪声档:差 2~6 倍", color=C_ALT, fontsize=9) fig.tight_layout() fig.savefig(os.path.join(FIGDIR, "excess_error.png"), dpi=130) plt.close(fig) # ───────────────────── 图 3:每步注入的噪声账本 ───────────────────── def fig_variance(sc, mix): s2_exact = LAB.sigma2_exact_posterior(sc) tt = np.arange(1, T + 1) tr_exact = np.array([np.trace(s2_exact[t]) / D for t in range(1, T + 1)]) fig, axes = plt.subplots(1, 2, figsize=(13.2, 5.0)) ax = axes[0] ax.semilogy(tt, sc.bt[1:], color=C_MAIN, lw=2.0, label=r"$\tilde\beta_t$(DDPM 反向链用的,fixed_small)") ax.semilogy(tt, sc.beta, color=C_ALT, lw=1.8, ls="--", label=r"$\beta_t$(fixed_large)") ax.semilogy(tt, tr_exact, color=C_GREEN, lw=2.0, label=r"真实后验方差($\tilde\beta_t$ + 均值那一项的不确定性)") ax.set_xlabel(r"时间步 $t$") ax.set_ylabel("每步注入噪声的方差") ax.set_title(r"(a) 真实后验方差比 $\tilde\beta_t$ 大,差值就是 NLL 缺口") ax.set_ylim(1e-6, 1e0) ax.grid(alpha=0.25, which="both") ax.legend(fontsize=9) ax = axes[1] lam = np.asarray(mix["lam"]) sr = np.asarray(mix["std_ratio"]) nn = np.asarray(mix["nll_uniform"]) ax.plot(lam, sr, "o-", color=C_MAIN, lw=2.0, label="生成分布的标准差 / 真分布") ax.axhline(1.0, color=C_GRAY, lw=1.2, ls=":") ax.set_xlabel(r"插值系数 $\lambda$:$\sigma^2=(1-\lambda)\tilde\beta_t+\lambda\beta_t$") ax.set_ylabel("标准差之比(1 表示胖瘦刚好对上)") ax.set_ylim(0.960, 0.978) ax2 = ax.twinx() ax2.plot(lam, nn, "s--", color=C_ALT, lw=1.8, label="期望 NLL") ax2.set_ylabel("期望 NLL(nats)", color=C_ALT) ax2.tick_params(axis="y", labelcolor=C_ALT) ax.set_title(r"(b) 换成 $\beta_t$ 把 3.4% 的偏窄补回 0.6 个点(没补满)", fontsize=11) ax.grid(alpha=0.25) fig.tight_layout() fig.savefig(os.path.join(FIGDIR, "variance_ledger.png"), dpi=130) plt.close(fig) # ──────────────────── 图 4:反向一步的三个配料 ──────────────────── def fig_step(sc): tt = np.arange(1, T + 1) abar_prev = sc.abar[:-1] abar_cur = sc.abar[1:] beta = sc.beta # x_{t-1} = c_x * x_t + c_0 * x_hat0 + sigma_t * z c_x = np.sqrt(sc.alpha) * (1.0 - abar_prev) / (1.0 - abar_cur) c_0 = np.sqrt(abar_prev) * beta / (1.0 - abar_cur) sigma = np.sqrt(sc.bt[1:]) fig, ax = plt.subplots(figsize=(9.2, 5.4)) ax.semilogy(tt, c_0, color=C_MAIN, lw=2.2, label=r"拉向 $\hat x_0$ 的系数(这一步挪了多远)") ax.semilogy(tt, c_x, color=C_GRAY, lw=1.8, ls="--", label=r"保留 $x_t$ 的系数") ax.semilogy(tt, sigma, color=C_ALT, lw=2.0, label=r"注入噪声的标准差 $\sqrt{\tilde\beta_t}$") ax.set_xlabel(r"时间步 $t$") ax.set_ylabel("系数 / 标准差(绝对值)") ax.set_title("反向一步的三个配料:绝大多数步子只挪千分之几") ax.grid(alpha=0.25, which="both") ax.legend(fontsize=9) ax.set_ylim(1e-6, 3.0) ax.annotate(rf"$t=500$:只挪 {c_0[499]:.2e}", xy=(500, c_0[499]), xytext=(560, 2e-5), fontsize=9, color=C_MAIN, arrowprops=dict(arrowstyle="->", color=C_MAIN, lw=1.2)) ax.annotate(rf"$t=2$:一步挪 {c_0[1]:.2f}", xy=(2, c_0[1]), xytext=(90, 0.9), fontsize=9, color=C_MAIN, arrowprops=dict(arrowstyle="->", color=C_MAIN, lw=1.2)) fig.tight_layout() fig.savefig(os.path.join(FIGDIR, "step_anatomy.png"), dpi=130) plt.close(fig) def main(): sc = LAB.Sched(LAB.linear_beta(T)) w_tilde = sc.weight("tilde") w_large = sc.weight("large") grid = list(range(1, T + 1)) floor = np.array([LAB.mse_floor(sc, t) for t in grid]) learn = np.array([LAB.learnability(sc, t) for t in grid]) print(" 重新训练三份闭式解权重(与 ddpm_lab.py 同一套种子)...") W, _ = LAB.section_B(sc, w_tilde, w_large) prof = {} for name, Wm in W.items(): prof[name] = np.array([ LAB.mse_analytic(*LAB.model_Bc(Wm, LAB.psi_basis(np.array([t]))[0]), sc, t) for t in grid]) mix = LAB.section_E(sc, W) fig_weight_profile(sc, w_tilde, w_large, learn) fig_excess(sc, grid, prof, floor, learn) fig_variance(sc, mix) fig_step(sc) print(" 四张图已写入 figures/") return dict(sc=sc, W=W, prof=prof, floor=floor, learn=learn, mix=mix) if __name__ == "__main__": main()
2026年09月27日
5 阅读
0 评论
0 点赞
2026-09-27
AIGC 基本功|扩散过程的前向与反向推导-SDE
扩散过程的前向与反向推导 所属方向:数学基础 | 难度:入门 | 前置知识:变分下界与重参数化(KL、重参数化技巧、重参数化梯度) 关键词:马尔可夫链、前向扩散、反向去噪、随机微分方程、score matching、DDPM、DDIM、DPM-Solver 01. 为什么需要它 先摆四组数字,全部来自文末附录里五个能直接跑的脚本。 第一组:同一个 score 下比较不同采样器。 在 2D 七分量高斯混合上,score 有闭式解,可以排除神经网络拟合误差。下表记录一次固定种子实验的 dlogp,即生成样本与参考样本的平均 $\log p_0(x)$ 之差。它只是一个分布统计量:0 只说明这一个均值相同,不代表分布完全相同;正值也不是“生成得更好”。 配置 score 评估次数 dlogp DDPM 祖采样 N=1000 1000 −0.032 DDPM 祖采样 N=200 200 +0.089 DDIM(η=0)N=50 50 +0.090 Heun 二阶 N=50 100 +0.019 DDIM@50 与 DDPM@200 在这次运行中的 dlogp 接近,但不能据此得出通用的“4 倍提速”。Heun@50 的 +0.019 与 DDPM@1000 的 −0.032 也不足以证明谁更准确:还需重复随机种子、估计差值的不确定性,并结合模式占比和模式内半径等指标。比较成本时要按 score 调用次数 NFE 计,Heun 每步调用两次。 不把这件事算清楚,调 num_inference_steps 就是盲猜:既不知道收益的量级,也分不清收益里哪部分来自「换了算法的阶数」、哪部分来自「少走了冤枉路」。 第二组:调度改变了各噪声水平上的计算分配。 本实验把线性调度换成余弦,DDIM@50 的单次 dlogp 从 0.090 变到 0.050。线性调度有 74.0% 的离散步处在单位方差参考下 SNR<1 的区域,这描述了噪声水平分配,并不证明这些步骤“没有用”;高噪声阶段也负责形成全局结构。0.040 的差值需要重复实验确认,不能称为已经验证的 1.8 倍质量提升。 这张图要看什么:左图是信号方差系数,右图是在数据方差为 1 的约定下的信噪比。虚线标出 $\bar\alpha=0.5$ 的位置;低于这条线仍可能保留有用语义,不能把阴影区直接标成无效计算。 第三组:换个参数化(ε / v / x₀),等于给每个时间步换了 8 个数量级的权重。 三种参数化描述的是同一个量,但作为最小二乘目标并不等价。把它们的损失都折算回「对 ε 误差的权重」:ε 参数化恒为 1;v 参数化跨 2.5×10⁴ 倍;x₀ 参数化跨 2.5×10⁸ 倍(t=1 时 1.0×10⁻⁴,t=1000 时 2.5×10⁴)。这表明改变参数化会改变隐式时间权重。DDPM 的简化目标是经样本质量实验支持的重新加权,不能仅由此表推出其唯一理由。 第四组:验证手段本身有个坑。 我第一版是用 MMD² 给采样器排序的,跑出来的表看着非常漂亮:Heun 在 N≥10 就顶到「噪声地板」,其余采样器一路降到 0 以下。问题是那个「地板」是单次抽样的运气值——拿真实样本去对真实样本、重复 8 次,MMD² 的标准差是 ±3.6×10⁻⁴,比 N≥100 时各采样器之间的差别还大。发现它的办法很笨:做一次噪声标定。附录里的 mmd_noise_check.py 就是干这个的。分布距离的估计量本身有方差,用它排序前先量一下噪声。 02. 最小可用理解 三句话。 第一句:前向是一串手工设好的加噪,能一步算出来。 它形式上是一条 1000 步的马尔可夫链,但因为每步加的都是高斯,累乘之后仍然是高斯,所以 $q(x_t|x_0)$ 有闭式解,不需要真的跑 1000 次。 第二句:反向需要「学」的只有一个东西——score。 给定 $x_t$ 去猜刚才加了什么噪声,最小均方误差意义下的最优答案就是 $\nabla_{x_t}\log p_t(x_t)$ 的一个常数倍。把这层说清楚之后,「从噪声生成数据」就退化成一个纯粹的数值积分问题。 第三句:采样器既选择过程,也选择离散化。 DDPM / 反向 SDE、DDIM / 概率流 ODE 的随机性与漂移不同;求解器还要选择参数化、时间网格、阶数、方差和预测后处理,不能只把它们看作同一个更新式改几个系数。 这张图要看什么:左图看「抹掉」的过程——t=200 时七个模式已经开始互相渗透,t=500 彻底变成一个圆球,看不出原始结构;右图看「找回来」的过程——40 条从同一个噪声球出发的轨迹,在最后十几步才各自「决定」落进哪个模式,前面漫长的路程只是在搭粗轮廓。这解释了大步长为什么主要伤细节而不伤整体布局。 03. 数学推导 3.1 前向链:为什么能一步到位 设原始数据 $x_0 \sim p_0$。前向过程每一步只做一件事:往上加一点高斯噪声。 $$q(x_t|x_{t-1}) = \mathcal{N}(x_t; \sqrt{\alpha_t}\,x_{t-1},\ \beta_t I)$$ 这里 $\alpha_t = 1 - \beta_t$,而 $\beta_t$ 是方差而不是标准差——这个记号是 DDPM 原文定的,容易看错。写成重参数化形式就是 $x_t = \sqrt{\alpha_t}\,x_{t-1} + \sqrt{\beta_t}\,\varepsilon$。 现在把两步接起来看: $$x_t = \sqrt{\alpha_t\alpha_{t-1}}\,x_{t-2} + \sqrt{\alpha_t\beta_{t-1}}\,\varepsilon_1 + \sqrt{\beta_t}\,\varepsilon_2$$ 两个独立高斯的线性组合仍是高斯,噪声方差为 $(\alpha_t\beta_{t-1}+\beta_t)I$。代入 $\beta_{t-1}=1-\alpha_{t-1}$、$\beta_t=1-\alpha_t$,得到 $\alpha_t(1-\alpha_{t-1})+(1-\alpha_t)=1-\alpha_t\alpha_{t-1}$。这才是两步方差的正确化简。 于是归纳下去,任意步数都能一步算出来。记 $\bar\alpha_t = \prod_{s\le t}\alpha_s$: $$q(x_t|x_0) = \mathcal{N}\big(x_t;\ \sqrt{\bar\alpha_t}\,x_0,\ (1-\bar\alpha_t)I\big)$$ 这就是整个扩散模型里唯一一条「白送」的公式,训练时要多少步的加噪样本都能直接算。注意它成立的前提是噪声必须是各向同性高斯:换成重尾噪声,累乘就不再是同类分布,这条式子立刻失效。 附录里的 forward_diffusion.py 拿 20 万样本对着验了一遍:用「逐步迭代 1000 次」和「一步闭式」两条路分别算 $x_t$ 的均值与方差,误差量级都在 $10^{-3}$,正好是 20 万样本的蒙特卡洛噪声水平(t=1000 时闭式方差误差 6.20×10⁻³,逐步方差误差 6.19×10⁻³,两者几乎相等,说明闭式解没错)。 3.2 反向链:真实后验长什么样 我们想要的是 $q(x_{t-1}|x_t)$,它不好算。但加上 $x_0$ 之后就好算了——贝叶斯公式、配方、得到 $$q(x_{t-1}|x_t,x_0) = \mathcal{N}\Big(x_{t-1};\ \tilde\mu_t,\ \tilde\beta_t I\Big)$$ 其中 $\tilde\beta_t = \dfrac{1-\bar\alpha_{t-1}}{1-\bar\alpha_t}\beta_t$,$\tilde\mu_t = \dfrac{\sqrt{\bar\alpha_{t-1}}\,\beta_t}{1-\bar\alpha_t}x_0 + \dfrac{\sqrt{\alpha_t}(1-\bar\alpha_{t-1})}{1-\bar\alpha_t}x_t$。 $\tilde\beta_t$ 是在已知 $x_0$ 和 $x_t$ 后的剩余不确定性。$t=1$ 时 $\bar\alpha_0=1$,所以 $\tilde\beta_1=0$;早期分子分母之比可能明显小于 1,后期二者都接近 1 时才有 $\tilde\beta_t\approx\beta_t$。它不是只在中间段才变小。 问题在于 $\tilde\mu_t$ 里含着 $x_0$,而这个量正是我们不知道的。DDPM 的做法是让网络去猜:把 $\tilde\mu_t$ 里的 $x_0$ 换成一个网络估计 $\hat x_0$,就得到一个可采样的反向链。 3.3 从 ELBO 到「只预测噪声」 把前向过程 $q(x_{1:T}\mid x_0)$ 作为变分分布,生成模型使用反向链 $p_\theta(x_{0:T})$,负 ELBO 分解为: $L_T=\mathrm{KL}(q(x_T\mid x_0)\Vert p(x_T))$——条件终端先验项。它依赖数据与固定调度,不依赖去噪网络。对本文数据取期望,修正脚本的闭式结果为线性调度 $1.065154\times10^{-4}$ nats、余弦调度 $6.410060\times10^{-9}$ nats。原先蒙特卡洛测的是边缘 $\mathrm{KL}(q(x_T)\Vert p(x_T))$,两者相差 $I(x_0;x_T)$,不能混用。忽略固定 $L_T$ 不改变网络梯度,但报告似然下界时仍需计入。 $L_{t-1} = \mathrm{KL}\big(q(x_{t-1}|x_t,x_0)\,\|\,\text{反向链}\big)$——去噪匹配项,起主导作用。 $L_0 = -\log p_\theta(x_0|x_1)$——最终重建/解码项,离散像素需要相应离散化似然;连续数据也不能不经评估就认定该项很小。 主项是两个高斯之间的 KL。两个同协方差高斯 $p=\mathcal{N}(\mu_p,\Sigma),q=\mathcal{N}(\mu_q,\Sigma)$ 的 KL 正好是 $\tfrac12(\mu_p-\mu_q)^\top\Sigma^{-1}(\mu_p-\mu_q)$——只剩均值之间的加权距离,协方差项被消掉了。代入 $\tilde\mu_t$ 和它的网络版本,需要网络拟合的只有 $\hat x_0$(或者等价地,$\varepsilon$)这一个量。 接下来是最关键的一次代入。用 $x_t = \sqrt{\bar\alpha_t}x_0 + \sqrt{1-\bar\alpha_t}\,\varepsilon$ 把 $x_0$ 换成 $\varepsilon$,会得到 $$\tilde\mu_t = \frac{1}{\sqrt{\alpha_t}}\Big(x_t - \frac{\beta_t}{\sqrt{1-\bar\alpha_t}}\,\varepsilon\Big)$$ 也就是均值里对 $\varepsilon$ 的依赖是线性的、系数是确定的。既然输出只需选一种等价参数化(每个数据坐标仍有一个分量),那就干脆让网络直接输出 $\varepsilon$,损失变成 $$L_{\text{simple}} = \mathbb{E}_{t,x_0,\varepsilon}\Big[\big\|\varepsilon - \varepsilon_\theta\big(\sqrt{\bar\alpha_t}x_0 + \sqrt{1-\bar\alpha_t}\,\varepsilon,\ t\big)\big\|^2\Big]$$ 从 ELBO 到这一行,中间扔掉了一串只跟 t 有关的系数。 为什么可以扔,06 节会用数字回答。 3.4 连续化:同一件事的两种时间写法 1000 步只是一个离散近似。把步长做成无穷小,$\beta_t$ 变成一个率 $\beta(u)$、$u\in[0,1]$,马尔可夫链就变成一个随机微分方程: $$dx = -\tfrac12\beta(u)\,x\,du + \sqrt{\beta(u)}\,dw$$ 这就是 Song 等人说的 VP-SDE(方差保持型)。系数 $-\tfrac12\beta(u)$ 让数据慢慢收缩到 0,$\sqrt{\beta(u)}$ 同时在加噪声——其边缘方差为 $\bar\alpha(u)\mathrm{Var}(x_0)+1-\bar\alpha(u)$:初始方差为 1 时才严格保持 1,否则逐渐趋向 1。 离散和连续能不能对上,是有条件的。 离散的 $\bar\alpha_t = \prod(1-\beta_s)$ 与连续的 $\exp(-\int_0^u \beta)$,只有在 $\beta$ 足够小时才接近。附录实测:线性调度折算回连续时间之后 $\beta_{\min}=0.1$、$\beta_{\max}=20$(正好是 Song 等人论文里的默认值),两条路的相对差从 $u=0.25$ 的 7.74×10⁻⁴ 涨到 $u=1$ 的 7.01×10⁻²。也就是说,连续时间的那套结论不能无条件搬到离散实现上——这个 7% 就是「离散化误差」的本体。 正向 SDE 对应一个 Fokker–Planck 方程,描述整个分布 $p_u$ 怎么随时间流动。关键观察:同一个 Fokker–Planck 方程对应无穷多条 SDE,它们的漂移项不同、但边缘分布完全相同。其中两条特别有用: $$\text{反向 SDE:}\quad dx = \Big[-\tfrac12\beta x - \beta\,\nabla_x\log p_u(x)\Big]du + \sqrt{\beta}\,d\bar w$$ $$\text{概率流 ODE:}\quad dx = \Big[-\tfrac12\beta x - \tfrac12\beta\,\nabla_x\log p_u(x)\Big]du$$ 这两条式子之间只差两个地方:score 前面的系数,和那个噪声项。 系数从 1 减半到 1/2,删掉噪声——因为注噪声带来的那部分扩散,正好被「一半的 score」抵消了,两者合起来保持边缘分布不变。所有采样器都在这两条式子之间做取舍,这就是 02 节第三句话的出处。 一个我实际踩的符号坑。 正向时间 $u$ 从 0 涨到 1,反向采样是让 $u$ 往回走,所以 $du<0$。我第一版把这一步忘了,写成 $x \leftarrow x - \tfrac12\beta(x+s)$,结果 Euler-Maruyama、概率流 ODE、Heun 三个采样器全线崩掉(MMD² 卡在 0.5 下不来,而 DDPM/DDIM 正常)。正确的写法是把「每步跨过的积分量」先抠出来: $$L = \int_{u_{\text{prev}}}^{u_{\text{cur}}}\beta(u)\,du = \log\frac{\bar\alpha_{\text{prev}}}{\bar\alpha_{\text{cur}}} > 0$$ 代入 $du=-1/N$($N$ 是采样步数)之后符号整体翻转,Euler 步变成 $x \leftarrow x + L\,(\tfrac12 x + c\,s)$,其中 $c=1$ 是反向 SDE、$c=1/2$ 是概率流 ODE。物理上很好理解:反向过程是把被前向压扁的分布吹回原样,漂移当然要往外推。 3.5 四种参数化:同一个量的四个名字 代码里同一个东西有四种写法,它们之间全是恒等式: $$\varepsilon = -\sqrt{1-\bar\alpha_t}\ \nabla_x\log p_t(x)$$ $$\mathbb{E}[x_0|x_t] = \frac{x_t + (1-\bar\alpha_t)\nabla_x\log p_t(x_t)}{\sqrt{\bar\alpha_t}}$$ 第二条就是 Tweedie 公式,它说的是「去噪」和「算 score」是同一件事。第一条则把 score 和 ε 预测对上。把它代进第二条,就得到代码里那句最眼熟的 pred_x0 = (x - sqrt(1-a) * eps) / sqrt(a)。 第四种是 v 参数化:$v = \sqrt{\bar\alpha_t}\,\varepsilon - \sqrt{1-\bar\alpha_t}\,x_0$。它看起来像个随手拼出来的组合,实际上 $(x_t, v)$ 和 $(x_0, \varepsilon)$ 之间是一个旋转: $$x_0 = \sqrt{\bar\alpha_t}\,x_t - \sqrt{1-\bar\alpha_t}\,v,\qquad \varepsilon = \sqrt{1-\bar\alpha_t}\,x_t + \sqrt{\bar\alpha_t}\,v$$ 验证很简单:把 $x_t = \sqrt{\bar\alpha_t}x_0 + \sqrt{1-\bar\alpha_t}\varepsilon$ 代入第一个式子,$\sqrt{\bar\alpha_t}\sqrt{1-\bar\alpha_t}$ 的交叉项正好抵消,剩下 $(\bar\alpha_t + 1 - \bar\alpha_t)x_0 = x_0$。变换矩阵的行列式是 $\bar\alpha_t + (1-\bar\alpha_t)=1$,正交——这就是 v 参数化「不会放大噪声」的来源。 score_bridges.py 把这四条恒等式全核了一遍:Tweedie 公式与精确后验均值的最大绝对误差在 $10^{-15}\sim10^{-13}$(t=1 时 1.332×10⁻¹⁵,t=1000 时 1.550×10⁻¹³),ε 的两条算法误差在 $10^{-16}\sim10^{-15}$,v 的重建误差 8.882×10⁻¹⁶。另外还做了一次不依赖闭式解的交叉验证:用 100 万样本的重要性采样直接估 $\mathbb{E}[x_0|x_t]$,与 Tweedie 公式的答案在小数点后两到三位一致(有效样本数从 t=50 时的 49895 涨到 t=1000 时的 999852)。 04. 代码实现 4.1 实验设计:把「模型误差」这个变量消掉 真实扩散模型的采样误差来自两处:score 估计得不准,以及数值积分不准。想把第二处单独看清楚,就得让第一处等于零——用一个 score 有闭式解的目标分布。 办法是选高斯混合:$p_0 = \sum_k w_k\mathcal{N}(\mu_k, s_k^2 I)$。前向加噪之后,第 $k$ 个分量的均值缩到 $\sqrt{\bar\alpha_t}\mu_k$、方差变成 $\bar\alpha_t s_k^2 + (1-\bar\alpha_t)$,仍然是各向同性高斯,所以 $p_t$ 还是高斯混合,且 $$\nabla_x\log p_t(x) = -\sum_k r_k(x)\frac{x - \sqrt{\bar\alpha_t}\mu_k}{\bar\alpha_t s_k^2 + (1-\bar\alpha_t)},\qquad r_k(x) = \frac{w_k\mathcal{N}_k(x)}{\sum_j w_j\mathcal{N}_j(x)}$$ $r_k$ 就是「这个样本属于第 $k$ 个分量」的责任度,是标准的软分配。于是整条反向链路都可以用真值跑,跑出来的差异 100% 来自离散化和「要不要注噪声」。 这里有个调参坑值得记一下:我第一版用了 3 个很宽的分量(标准差 0.45/0.35/0.55),结果五种采样器全部顶到 MMD 噪声地板上,分不出高下。换成 7 个紧分量(标准差 0.16~0.30)之后差距才显出来——目标太光滑,分辨不出采样器的差别。这和真实情况是对应的:scheduler 之间的差别本来就在高频细节上。 4.2 前向:闭式解核对 核心就几行,除了算 $\bar\alpha$ 之外没有任何魔法: def linear_beta(T=1000, beta_1=1e-4, beta_T=0.02): return np.linspace(beta_1, beta_T, T) def alpha_bar_from_beta(beta): return np.cumprod(1.0 - beta) def q_sample(x0, t_idx, alpha_bar, rng): a = alpha_bar[t_idx] eps = rng.standard_normal(x0.shape) return np.sqrt(a) * x0 + np.sqrt(1.0 - a) * eps, eps q_sample 就是 3.1 节那个闭式解。余弦调度多两行——按 $\bar\alpha_t$ 定义式算完之后要把 β 截到 0.999 以内,因为 $t=T$ 时 $\cos(\pi/2)=0$ 会让最后一步 $\beta=1$,实操直接炸: def cosine_beta(T=1000, s=0.008, clip=0.999): t = np.arange(1, T + 1) / T f = np.cos(((t + s) / (1 + s)) * np.pi / 2) ** 2 f0 = (np.cos((s / (1 + s)) * np.pi / 2)) ** 2 beta = beta_from_alpha_bar(f / f0) # 反解出每步的 beta return np.clip(beta, None, clip) 4.3 反向:五个采样器 DDPM 祖采样——用高斯核近似反向转移。对一般数据,即使均值来自精确 score,有限步高斯核也不等于真实反向条件分布。把 3.2 节的 $\tilde\mu_t$ 和 $\tilde\beta_t$ 抄进来,再把 $\hat x_0$ 换成 Tweedie 公式: def run_ddpm(x, alpha_bar, taus, rng): for j in range(len(taus) - 1): t_cur, t_prev = taus[j], taus[j + 1] a_cur, a_prev = abar_at(alpha_bar, t_cur), abar_at(alpha_bar, t_prev) alpha_j = a_cur / a_prev # 跨 k 步的等效 alpha beta_j = 1.0 - alpha_j s = score_at(x, alpha_bar, t_cur) mean = (x + beta_j * s) / np.sqrt(alpha_j) var = beta_j * (1.0 - a_prev) / (1.0 - a_cur) x = mean + np.sqrt(var) * rng.standard_normal(x.shape) return x Euler-Maruyama(反向 SDE)——把 $L$ 和 $c=1$ 代进 3.4 节的结果: def run_em(x, alpha_bar, taus, rng): for j in range(len(taus) - 1): L = _log_step(alpha_bar, taus[j], taus[j + 1]) s = score_at(x, alpha_bar, taus[j]) x = x + L * (0.5 * x + s) + np.sqrt(L) * rng.standard_normal(x.shape) return x 概率流 ODE——唯一的改动是 score 系数减半、噪声项删掉: def run_ode(x, alpha_bar, taus, rng): for j in range(len(taus) - 1): L = _log_step(alpha_bar, taus[j], taus[j + 1]) s = score_at(x, alpha_bar, taus[j]) x = x + 0.5 * L * (x + s) return x Heun 二阶——Euler 预测一步,再用终点的 score 校正一次。每步两次评估: def run_heun(x, alpha_bar, taus, rng): for j in range(len(taus) - 1): L = _log_step(alpha_bar, taus[j], taus[j + 1]) s0 = score_at(x, alpha_bar, taus[j]) d0 = 0.5 * L * (x + s0) x1 = x + d0 # Euler 预测 s1 = score_at(x1, alpha_bar, taus[j + 1]) # 终点再评估一次 d1 = 0.5 * L * (x1 + s1) x = x + 0.5 * (d0 + d1) return x DDIM 单独说,因为它长得不像上面四个。它先把样本一步跳到 $\hat x_0$,再按目标时刻的 $\bar\alpha$ 重新加回噪声: $$x_{t-1} = \sqrt{\bar\alpha_{t-1}}\,\hat x_0 + \sqrt{1-\bar\alpha_{t-1}}\,\hat\varepsilon$$ $\hat x_0$ 用 Tweedie 公式算、$\hat\varepsilon$ 用 $\varepsilon=-\sqrt{1-\bar\alpha_t}s$ 算。η=0 的 DDIM 在小步长极限对应概率流 ODE,其 score 漂移系数为反向 SDE 的一半;DDPM 的均值更新对应后者,两者并不在代数上相同。 4.4 结果 七个分量、每档生成 4000 个样本、参照集 8000 个真实样本。三个主指标都先扣掉了「真实样本自己」的基线,所以 0 才等于完美: 采样器 N=10 N=25 N=50 N=100 N=200 N=1000 DDPM −0.365 +0.266 +0.249 +0.174 +0.089 −0.032 DDIM −0.248 +0.124 +0.090 +0.057 +0.032 +0.007 EMA(SDE) −2.725 −0.756 −0.347 −0.189 −0.100 −0.081 Euler(ODE) −2.340 −0.792 −0.348 −0.159 −0.074 −0.014 Heun +0.546 +0.084 +0.019 +0.005 +0.002 +0.001 上表是单次实验的 dlogp。两批独立真实样本在这次运行中的差值约 0.053;这不是经重复估计的标准差或置信区间。表格主要帮助识别量级很大的离散化偏差,微小差异不做显著性排序。三点读法: DDPM / DDIM 在 N=10 还能看,Euler 系直接崩(−2.3 ~ −2.7,说明样本散在低密度区根本没收敛)。这不能怪 ODE 或 SDE,只能怪第一步就跨了 1.92 的积分量。 误差变号这件事有意义。 DDPM/DDIM 的 dlogp 是正的(样本被堆到高密度区、偏聚拢),Euler 系是负的(还没收敛、偏散开)。两种失败模式方向相反,光看一个「距离」指标看不出来。 本次 Heun 在 N≥20 的 dlogp 较小,但差异接近采样波动时不作排名;接近0也不足以证明整个分布正确。 补充两个不同统计量:按模式归属估算的分量占比 TV,以及模式内均方半径的相对误差。在 N=50 的本次运行中,Heun 为 0.011 / 0.066,DDIM 为 0.014 / 0.155,DDPM 为 0.019 / 0.340,Euler-ODE 为 0.017 / 0.358。它们分别检查模式质量与分散程度,排序并非处处一致,也要估计采样误差。 这张图要看什么:左图画的是单个诊断统计量的偏差,参考虚线只表示一次真实样本对照差值,不是置信区间。右图比较线性漂移项的指数因子和 Euler 近似,解释大步长为何可能困难;完整 score 场同时参与更新,不能仅用这一项定量归因全部误差。 4.5 为什么朴素 Euler 在大步长下会崩 把上面第 1 点挖到底。概率流 ODE 里线性部分 $\tfrac12\beta x$ 的精确解是 $e^{L/2}$,而 Euler 用的是它的一阶展开 $1+L/2$。两者的相对误差随 $L$ 指数上升: 采样步数 N 每步最大积分量 $L_{\max}$ $e^{L/2}$ 与 $1+L/2$ 的相对误差 10 1.9197 24.95% 20 0.9852 8.80% 50 0.4002 1.75% 100 0.2011 0.47% 200 0.1008 0.12% 1000 0.0202 0.01% DDPM 和 DDIM 把这部分精确解掉了——它们的更新直接把 $\sqrt{\bar\alpha}$ 乘上去,等于用指数积分器而不是 Euler。所以 N=10 时它们还能看,而 Euler 每步有 25% 的相对误差、误差还会沿着链累积。这也解释了为什么工业实现里大家都用 DDIM/PNDM/DPM-Solver 的写法,而不是拿反向 SDE 直接上 Euler。 高阶方法也需要合适的步长。 本次 Heun@10 的 dlogp 为 +0.546,DDIM@10 为 −0.248,说明这个网格上的 Heun 出现明显偏聚拢。原因可能包括稳定域、非线性 score 与步长的相互作用;不能说二阶校正必然放大线性项误差,更不能推广为 DPM-Solver 高阶算法在所有少步数任务上都更差。 05. 工业级实现对照 对着 huggingface/diffusers 的 src/diffusers/schedulers/scheduling_ddpm.py 看(https://github.com/huggingface/diffusers/blob/main/src/diffusers/schedulers/scheduling_ddpm.py,以 2026-09 时的实现为准)。 add_noise 就是 3.1 节的闭式解,一字不差: sqrt_alpha_prod = alphas_cumprod[timesteps] ** 0.5 sqrt_one_minus_alpha_prod = (1 - alphas_cumprod[timesteps]) ** 0.5 noisy_samples = sqrt_alpha_prod * original_samples + sqrt_one_minus_alpha_prod * noise 第一个值得注意的细节:noise 是从外面传进来的,不是在函数里采的。因为训练时需要知道「这次加的是哪个 ε」才能算损失——如果函数内部自己采样,你就永远拿不到那个 target。 第二个细节:alphas_cumprod 在 __init__ 里用 torch.cumprod(1 - betas) 一次算好并缓存。1000 个数,每个训练步都要按 timestep 取,重新累乘显然不划算。这就是 3.1 节那条「白送」的公式在工程上的直接体现。 step 与我的 numpy 实现是代数等价的。 它写的是 DDPM 原文公式 (7) 的系数形式: current_alpha_t = alpha_prod_t / alpha_prod_t_prev current_beta_t = 1 - current_alpha_t pred_original_sample = (sample - beta_prod_t ** 0.5 * model_output) / alpha_prod_t ** 0.5 # 即 Tweedie pred_original_sample_coeff = (alpha_prod_t_prev ** 0.5 * current_beta_t) / beta_prod_t current_sample_coeff = current_alpha_t ** 0.5 * beta_prod_t_prev / beta_prod_t pred_prev_sample = pred_original_sample_coeff * pred_original_sample + current_sample_coeff * sample 把 pred_original_sample_coeff 和 current_sample_coeff 代进 3.2 节的 $\tilde\mu_t$ 表达式展开,$x_t$ 的系数会化简成 $1/\sqrt{\alpha_j}$、$\varepsilon$ 的系数化简成 $-\beta_j/(\sqrt{\alpha_j}\sqrt{1-\bar\alpha_{\text{cur}}})$,也就是 $\tilde\mu = (x_t + \beta_j s)/\sqrt{\alpha_j}$——和我 4.3 节那个「一行 DDPM」完全一样。current_alpha_t = alpha_prod_t / alpha_prod_t_prev 对应的正是我的 alpha_j,所以 diffusers 支持跳步(strided)采样;这里只在相同 prediction_type、方差与不启用 clipping/thresholding 的条件下和最小实现等价。 三处「最小实现没有、工业实现必须有」的差异: 第一,调度不是自由参数,是分类的。beta_schedule 有 linear / scaled_linear / squaredcos_cap_v2 / sigmoid 几个分支;我 3.1 节实现的余弦调度在它这里叫 squaredcos_cap_v2,走的是通用函数 betas_for_alpha_bar——先给定 $\bar\alpha(t)$ 的解析式,再按 $1 - \bar\alpha(t_2)/\bar\alpha(t_1)$ 反解出每步 β,并且硬编码 max_beta=0.999。这正好对上我 cosine_beta 里那句 clip,不是巧合:$t=T$ 时 $\cos(\pi/2)=0$ 会给出 $\beta_T=1$,必须截。 第二,有限终端 SNR 会引入起点分布近似。线性表的 $\bar\alpha_T=4.0358\times10^{-5}$,仍保留少量数据成分;推理却常从标准高斯开始。diffusers 的 rescale_betas_zero_snr 对应 Lin 等,2023 的修正,但要与训练的参数化、时刻采样和起点设置一起核对。对已有 epsilon checkpoint 不能只开一个开关就假定兼容,零 SNR 下若仍用除以 $\sqrt{\bar\alpha_T}$ 的公式还会遇到奇异点。 第三,参数化是可切换的输出头。prediction_type 支持 epsilon / sample / v_prediction 三选一,step 开头那个 if-elif 就干这件事。三种模式对 pred_original_sample 的算法不同,但后面的系数计算完全共用——这正是 3.5 节「它们描述同一个量」在工程上的样子。也正因如此,采样端转换形式可由配置选择,但 checkpoint 的训练目标必须匹配,不能把 epsilon 权重仅改一行配置就当作 v 预测器,代价藏在 06 节的权重表里。 06. 代价与边界 代价一:换参数化不是改记号,是改损失权重。 把三种参数化都折算回「对 $\varepsilon$ 误差的权重」(推导见 score_bridges.py 的文档串:x₀ 的损失 $\|\delta x_0\|^2$ 折成 $\varepsilon$ 误差要乘倍率,v 的误差因为 $\delta\varepsilon = \sqrt{\bar\alpha_t}\,\delta v$ 也要乘): t $\bar\alpha_t$ ε 参数化的权重 v 参数化的权重 x₀ 参数化的权重 1 9.999×10⁻¹ 1.0000 1.0001 1.0001×10⁻⁴ 100 8.970×10⁻¹ 1.0000 1.1148 1.1480×10⁻¹ 500 7.859×10⁻² 1.0000 1.2725×10¹ 1.1725×10¹ 900 2.752×10⁻⁴ 1.0000 3.6336×10³ 3.6326×10³ 1000 4.036×10⁻⁵ 1.0000 2.4778×10⁴ 2.4777×10⁴ 这张表是等价残差之间的代数权重,不是网络参数梯度的实测值。 ε、v、x₀ 三种目标的输出尺度、网络雅可比以及时间采样方式都会影响梯度。等权 ε 损失在这里对应 ε 残差权重恒为 1;不能据此断言各时刻梯度相同,也不能仅凭该表断言 x₀ 目标一定被高噪声支配或无法训练。 这张图要看什么:把不同输出误差换算到 ε 误差时,各自带上不同时间权重。曲线说明训练目标不等价,不是实测梯度图;实际训练还要连同参数化和时间采样权重一起分析。 代价二:阶数更高也要选择合适的网格。 Heun@10 的表现说明本实验的大步长不合适,不是对所有高阶求解器的否定。DPM-Solver 专门利用扩散 ODE 的半线性结构,不能从朴素 Heun 的结果推导它的少步数表现。 代价三:随机项改变有限步采样的误差与方差。 DDIM 的 η 从 0 到 1,本次 N=50 的 dlogp 为 0.090、0.065、0.092、0.169、0.249。这里 η=1 配上相同时间网格和方差选择可恢复所实现的 DDPM 更新,但不能把单次结果概括为“噪声没有用”。在精确 score、连续时间和正确初始分布下,反向 SDE 与概率流 ODE 都能得到相同边缘分布;有限步误差的优劣取决于具体离散化。 边界:这套实验测不到模型误差。 oracle score 排除了训练误差,却仍有有限终端分布近似、离散化和样本统计误差;它不是现实图像模型误差的严格下界。真实网络的误差还会与采样轨迹交互,必须在目标 checkpoint 上另做实验。 边界:高斯混合不是自然图像。 这里的数据维数、模态结构、score 光滑性和引导强度都很简单。表格适合验证公式与数值方法,不能当作真实模型加速倍率或通用采样器排名。 ODE 也能生成多样样本与估计似然。 确定性只意味着固定初始噪声对应固定轨迹;不同随机初值仍可覆盖整个数据分布。原始 score-SDE 论文 就利用概率流 ODE 计算似然。在向量场满足正则条件时精确流可逆,但有限步 DDIM / ODE 数值求解一般不能无误差反演。 按步数预算给一张速查表(数字全部来自 4.4 节那张表,dlogp,绝对值越小越好): 步数预算 该选谁 实测依据 8~10 步 在目标模型上比较合适的少步采样器 本 toy 的 Heun@10 偏聚拢,不能推出其他高阶方法的排名 20~50 步 Heun 是本实验可考虑的方案 Heun@50 的 dlogp 为 0.019,需结合其他指标和重复试验 100~200 步 同 NFE 比较,避免只按步数选 N=200:Heun 0.002、DDIM 0.032、DDPM 0.089 要多样性 ODE / SDE 都可,比较分布覆盖 ODE 的多样性来自随机初值;SDE 还增加路径随机性 要反演 / 要编辑 可考虑概率流 ODE / DDIM 理想流可逆;有限步反演仍有数值和模型误差 表格只总结当前教学实验能支持的选择。没有实测的 PNDM / DPM-Solver 不参与排名;真实模型应以同等 NFE、重复种子与多个质量指标比较。 07. 经典论文脉络 ① 1503.03585(Sohl-Dickstein et al., 2015)——把扩散搬进生成模型。 用非平衡热力学里的一个想法:先定义一个把数据逐步破坏成噪声的正向过程,再学它的反向过程。贡献是框架本身,同时研究了高斯与二项扩散等设置,采样慢到没有实战价值。 ② 2006.11239(Ho et al., 2020)——DDPM,把目标改成「预测噪声」。 三件事:把 $\tilde\mu_t$ 参数化成预测 $\varepsilon$;指出把 ELBO 里那一串只跟 t 有关的系数扔掉、直接用等权的 $L_{\text{simple}}$ 反而效果更好;给出 3.2 节那套 $\tilde\beta_t / \tilde\mu_t$ 的闭式解。这才是「扩散模型能训练起来」的直接原因——在这之前,没人找到规模化的训练目标。 ③ 2011.13456(Song et al., 2021)——把离散和连续统一起来。 这篇是本节点的锚点。它做了两件大事:把 DDPM(VP-SDE)和它自己那套 score matching(VE-SDE)统一到同一个 SDE 框架下,把它们写成不同漂移与扩散系数下的 SDE,并由相应 Fokker–Planck 方程构造边缘等价的概率流 ODE;以及提出了概率流 ODE——同一个边缘分布、确定性求解、还能用现成的 ODE 求解器(这篇文章里用四阶 Runge–Kutta)。3.4 节那两条式子就出自这里。 ④ 2010.02502(Song et al., 2021)——DDIM,确定性采样。 把反向链改成不含随机项的确定性映射:样本轨迹只由 $x_T$ 决定,跳步采样不再需要「一步一小步」的假设。两篇作者不同:DDIM 的第一作者是 Jiaming Song,score-SDE 的第一作者是 Yang Song——DDIM 本质上就是概率流 ODE 的一个(指数积分器风格的)离散化。它让 50 步的采样第一次在质量上追平 1000 步。 ⑤ 2102.09672(Nichol & Dhariwal, 2021)——余弦调度。 指出线性调度把信噪比压得太快(就是 01 节第二组数字里那 74%),改成 $\bar\alpha_t = \cos^2(\cdot)$ 之后低步数下的质量明显更好。这篇的价值在于它把「调度」从工程细节变成了有图像解释的设计问题:$\bar\alpha_t$ 的形状决定了「每个时刻还剩多少信息量」,而余弦的形状让信息量的衰减更均匀。 在这五篇之外还有两条重要支线:v 参数化(arXiv:2202.00512,Salimans & Ho)解决 06 节权重表里 x₀ 参数化的尺度失衡问题;零终端 SNR(arXiv:2305.08891)修掉 05 节那个 $\bar\alpha_T\neq0$ 的 bug。它们都是在这条主线已经跑通之后,针对具体失效模式的补丁。 08. 常见误解 误解一:「DDPM 和 DDIM 是同一条 ODE 的两种离散化。」 DDPM 是带噪声的高斯反向马尔可夫链,连续极限对应反向 SDE;η=0 的 DDIM 是确定性路径,可联系到概率流 ODE。二者的连续边缘分布可一致,但漂移中的 score 系数不同:SDE 是 1,ODE 是 1/2,不能只把 DDPM 的噪声删除就得到 DDIM。 误解二:「反向过程就是把噪声一步步减掉。」 方向反了。反向 drift 是 $+\tfrac12\beta x + \beta s$,$x$ 那一项是往外推的。前向把分布压向原点,反向把它吹回原样。我按「减掉」实现了三个采样器,全部崩掉(4.3 节那个符号坑),而错误版本跑起来并不报错、只是数值不对——这类 bug 只能靠对着闭式解核对来抓。 误解三:「换个参数化只是记号问题,等价就是等价。」 描述的对象等价,作为训练目标不等价。权重表跨 8 个数量级(06 节),这不是小差异。等价只发生在「已经收敛到精确最优解」这个极限情况;训练过程中不同参数化走的路径完全不同。 误解四:「换成 mean log p 就能可靠排序。」 MMD / FID 需要估计不确定性,mean log p 同样需要,而且单个均值不能刻画完整分布。对两批真实样本只算一次差值,不能据此画置信区间;必须重复抽样或用适当的标准误 / bootstrap 分析。 误解五:「ODE 无随机项,所以没有多样性;SDE 多走几步只会累积坏噪声。」 随机性可以来自初始噪声,也可以来自路径。精确连续过程下两者都可得到正确分布;更多步数通常减少相应数值方法的离散误差,但有限模型、引导、网格与算力约束下仍应实测。 误解六:「$\bar\alpha_t$ 和 $1-\bar\alpha_t$ 加起来是 1,所以信噪比就是 $\bar\alpha_t$。」 信噪比是 $\bar\alpha_t/(1-\bar\alpha_t)$,不是 $\bar\alpha_t$。t=500 时线性调度 $\bar\alpha_{500}=7.86\times10^{-2}$,看着还有 7.9% 的信号,但信噪比只有 0.085(−10.7 dB)——信号和噪声的幅度比已经掉到 1:3.4 以下。只看 $\bar\alpha$ 会严重高估「还剩多少信息」。 09. 动手验证 五个脚本都在文末附录里,用 /usr/local/bin/python3 直接跑,只依赖 numpy 和 matplotlib(make_figures.py 需要 matplotlib,其余四个只要 numpy)。 python forward_diffusion.py # 前向闭式解核对 + 两种调度对比 + 终端 SNR python reverse_sampling.py # 五种采样器 × 七档步数的主实验 python score_bridges.py # Tweedie / ε / v 四条恒等式核对 + 权重表 python mmd_noise_check.py # 先量一下 MMD² 自己的噪声 python make_figures.py # 生成四张图 值得自己动手改着看的四处: 第一,把 N 设成 10,看 DDIM 和 Euler 谁先崩。 预期:DDIM 的 dlogp 是 −0.248(偏散、没收敛完),Euler 是 −2.340(散得离谱)。原因是 4.5 节那张表——N=10 时每步跨 1.92 的积分量,Euler 的线性部分误差 24.95%,而 DDIM 直接乘 $\sqrt{\bar\alpha}$ 把它精确解掉了。 第二,给 run_ode 加一个符号。 把 x + 0.5 * L * (x + s) 改成 x - 0.5 * L * (x + s) 再跑,预期:MMD² 从 8.57×10⁻⁴ 涨到 0.5 量级,dlogp 从 −0.348 变成毫无意义的正数。这就是 3.4 节那个 $du<0$ 的坑。 第三,把目标分布的 SCALES 从 0.16~0.24 改成 0.45 附近。 预期:五种采样器在 N≥50 时全部顶到噪声地板,表格失去分辨力。这是 4.1 节那个「目标太光滑就分不出来」的第二遍确认。 第四,改 forward_diffusion.py 里的 linear_beta,把 beta_T 从 0.02 降到 0.01。 预期:$\bar\alpha_T$ 变大(终端残留信号变多),t=500 的 SNR 也跟着变。可以顺着看 05 节说的那个「非零终端 SNR」到底有多敏感。 10. 延伸阅读 前置:变分下界与重参数化(vae_elbo)。3.3 节里「两个同协方差高斯的 KL 只剩均值距离」和「重参数化让梯度穿透采样」这两步,在前置那篇里有完整的推导和代码验证,这里直接用结论了。 后续:视频 VAE 的时空压缩结构(video_vae)。把 3.1 节的 2D 扩散扩到时空 $(T,H,W)$ 之后,噪声调度要不要沿时间轴做非均匀分配,是视频模型和图像模型第一个分岔点。 后续:FID / CLIP Score 到底测了什么(image_metrics)。08 节误解四的教训在通用指标上会放大——FID 的估计方差、样本量、以及它对「模式丢失」不敏感这件事,值得单独算一次。 后续:KV Cache 与自回归视频生成(kv_cache)。自回归视频模型把「一步去噪」换成「一次前向出一帧」,3.4 节的两条 SDE/ODE 就不再适用了,但 score 那一层的直觉仍然通用。 附录:完整代码 09 节用到的脚本全文如下(mmd_noise_check.py、forward_diffusion.py、score_bridges.py、make_figures.py、reverse_sampling.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 mmd_noise_check.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """MMD^2 到底能不能拿来给采样器排序?——结论:在这个问题上不能。 写这篇的时候第一版主指标就是 MMD^2,跑出来的表看着很漂亮(Heun 在 N>=10 就顶到"噪声地板",其余采样器数值一路降到 0 以下)。问题是那个 "噪声地板"是单次抽样的运气值:把真实样本对真实样本重复 8 次,MMD^2 的标准差是 3.6e-4,比 N>=100 时各采样器之间的差别还大。 所以正文改用三个有闭式解、方差小得多的指标(mean log p_0 / 分量占比 TV / 模式内半径),MMD^2 只留作低步数区间的旁证。 跑法:python mmd_noise_check.py """ import numpy as np from forward_diffusion import sample_data from reverse_sampling import mmd2, run_ddpm, make_stride from forward_diffusion import alpha_bar_from_beta, linear_beta, T, D ab = alpha_bar_from_beta(linear_beta(T)) print("A. 真实样本 vs 真实样本(应当 ~0),重复 8 次,看估计量的散布") for bw in [None, 1.9, 1.0, 0.5]: vals = [] for k in range(8): a = sample_data(2000, np.random.default_rng(1000 + k)) b = sample_data(2000, np.random.default_rng(5000 + k)) v, used = mmd2(a, b, bw=bw) vals.append(v) vals = np.array(vals) print(f" bw={str(bw):>5} (实取{used:.2f}) mean={vals.mean():+.3e} " f"std={vals.std():.3e} min={vals.min():+.3e} max={vals.max():+.3e}") print() print("B. ddpm@1000 换 8 个不同的初始噪声种子") for bw in [None, 1.0, 0.5]: ref = sample_data(4000, np.random.default_rng(7)) vals = [] for k in range(8): rng = np.random.default_rng(3000 + k) x = rng.standard_normal((2000, D)) x = run_ddpm(x, ab, make_stride(1000, T), rng) v, used = mmd2(x, ref, bw=bw) vals.append(v) vals = np.array(vals) print(f" bw={str(bw):>5} (实取{used:.2f}) mean={vals.mean():+.3e} " f"std={vals.std():.3e}") print(f" -> {np.array2string(vals, precision=3, formatter={'float': lambda v: f'{v:+.2e}'})}") print() print("C. 极端对照:把样本整体平移 0.5,看各带宽的分辨力") ref = sample_data(4000, np.random.default_rng(7)) for shift in [0.05, 0.1, 0.2, 0.5]: y = sample_data(2000, np.random.default_rng(11)) + shift row = f" shift={shift:.2f} " for bw in [1.9, 1.0, 0.5, 0.25]: v, _ = mmd2(y, ref, bw=bw) row += f"bw={bw}:{v:+.2e} " print(row) forward_diffusion.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """扩散前向过程:闭式解、噪声调度、信噪比。 只依赖 numpy。正文里出现的每一个数字都由本脚本打印,不手填。 要回答三个问题: 1. q(x_t | x_0) 的闭式解是不是真的成立(蒙特卡洛对着验) 2. 线性调度和余弦调度把"难度"分配得有多不一样(看 alpha_bar / SNR) 3. 前向终点的分布离标准正态到底差多少(决定了 L_T 那一项有多大) 运行: python forward_diffusion.py """ import numpy as np SEED = 20260926 T = 1000 # 总步数,与 DDPM 原文一致 D = 2 # 玩具数据维度 # ── 玩具数据:2D 七分量各向同性高斯混合(环状 + 一个中心分量)────────── # 分量故意取紧(标准差 0.16~0.30):score 场的曲率大,离散化误差才显出来。 # 换成 3 个宽分量(标准差 ~0.5)的话,oracle score 下所有采样器都会直接顶到 # MMD 噪声地板上,分不出高下——这是调这个玩具问题时踩的第一个坑。 WEIGHTS = np.array([0.14, 0.13, 0.15, 0.12, 0.14, 0.13, 0.19]) MEANS = np.array([[-2.6, -1.0], [-1.0, -2.2], [1.4, -2.0], [2.6, -0.4], [1.6, 1.8], [-0.6, 2.4], [0.0, 0.0]]) SCALES = np.array([0.22, 0.18, 0.20, 0.16, 0.24, 0.20, 0.30]) def sample_data(n, rng): """从 p_0 采 n 个样本。""" k = rng.choice(len(WEIGHTS), size=n, p=WEIGHTS) return MEANS[k] + SCALES[k][:, None] * rng.standard_normal((n, D)) # ── 噪声调度 ────────────────────────────────────────────────────────── def linear_beta(T=T, beta_1=1e-4, beta_T=0.02): """DDPM 原文的线性调度,beta 从 1e-4 均匀涨到 0.02。""" return np.linspace(beta_1, beta_T, T) def cosine_beta(T=T, s=0.008, clip=0.999): """Nichol & Dhariwal 的余弦调度,返回 beta_t。 alpha_bar_t = cos^2(((t/T + s)/(1+s)) * pi/2) / cos^2((s/(1+s)) * pi/2) 分母只是为了让 t=0 时 alpha_bar=1。 按 t=T 代入会得到 alpha_bar_T = 0(cos(pi/2)=0),也就是最后一步 beta=1, 实操上会炸。原文和 diffusers 都会把 beta 截到 0.999 以内,这里照做。 """ t = np.arange(1, T + 1) / T f = np.cos(((t + s) / (1 + s)) * np.pi / 2) ** 2 f0 = (np.cos((s / (1 + s)) * np.pi / 2)) ** 2 a_bar = f / f0 beta = beta_from_alpha_bar(a_bar) return np.clip(beta, None, clip) def cosine_alpha_bar(T=T, s=0.008, clip=0.999): """截尾之后的余弦调度的 alpha_bar,与 cosine_beta 一致。""" return alpha_bar_from_beta(cosine_beta(T, s, clip)) def alpha_bar_from_beta(beta): """alpha_bar_t = prod_{s<=t} (1 - beta_s),t = 1..T。""" return np.cumprod(1.0 - beta) def beta_from_alpha_bar(alpha_bar): """反解出每步的 beta_t = 1 - alpha_bar_t / alpha_bar_{t-1}。""" prev = np.concatenate([[1.0], alpha_bar[:-1]]) return 1.0 - alpha_bar / prev def snr_db(alpha_bar): """信噪比 alpha_bar/(1-alpha_bar),取 10*log10。""" return 10.0 * np.log10(alpha_bar / (1.0 - alpha_bar)) def q_sample(x0, t_idx, alpha_bar, rng): """x_t = sqrt(alpha_bar_t) * x_0 + sqrt(1 - alpha_bar_t) * eps。 t_idx 是 0-based 的数组下标,对应 alpha_bar[t_idx]。 """ a = alpha_bar[t_idx] eps = rng.standard_normal(x0.shape) return np.sqrt(a) * x0 + np.sqrt(1.0 - a) * eps, eps # ── 加噪后分布的解析形式(反向采样要用 oracle score)────────────────── def noised_mixture(alpha_bar_t): """p_t 仍是高斯混合:第 k 个分量的均值缩到 sqrt(a)*mu_k, 协方差变成 a*Sigma_k + (1-a)*I。因为 Sigma_k = s_k^2 I,结果仍是各向同性。 """ a = alpha_bar_t means = np.sqrt(a) * MEANS var = a * SCALES ** 2 + (1.0 - a) # 每个分量的方差(标量) return WEIGHTS, means, var def log_density(x, weights, means, var): """各向同性高斯混合的 log 密度,x: [n, 2]。""" n, d = x.shape # [n, K] sq = ((x[:, None, :] - means[None, :, :]) ** 2).sum(-1) comp = -0.5 * (sq / var[None, :] + d * np.log(2 * np.pi * var)[None, :]) mx = comp.max(axis=1, keepdims=True) e = np.exp(comp - mx) mix = (weights[None, :] * e).sum(1) return (np.log(mix) + mx[:, 0]) def score_fn(x, alpha_bar_t): """nabla_x log p_t(x) 的闭式解。混合权重用 log-sum-exp 稳住数值。""" w, m, v = noised_mixture(alpha_bar_t) n, d = x.shape sq = ((x[:, None, :] - m[None, :, :]) ** 2).sum(-1) comp = -0.5 * (sq / v[None, :] + d * np.log(2 * np.pi * v)[None, :]) mx = comp.max(axis=1, keepdims=True) e = np.exp(comp - mx) resp = w[None, :] * e resp = resp / resp.sum(1, keepdims=True) # [n, K] 责任度 # grad log N(x; m_k, v_k I) = -(x - m_k)/v_k return -(resp[:, :, None] * (x[:, None, :] - m[None, :, :]) / v[None, :, None]).sum(1) # ────────────────────────────────────────────────────────────────────── # A. 闭式解核对 # ────────────────────────────────────────────────────────────────────── def check_closed_form(rng, alpha_bar, n=200_000): """一步一步加噪 1000 次,和闭式解 x_t = sqrt(a) x_0 + sqrt(1-a) eps 对着验。""" x0 = sample_data(n, rng) out = [] for t_idx in [49, 199, 499, 999]: a = alpha_bar[t_idx] # 路径 1:逐步迭代 x = x0.copy() beta = beta_from_alpha_bar(alpha_bar) for i in range(t_idx + 1): x = np.sqrt(1.0 - beta[i]) * x + np.sqrt(beta[i]) * rng.standard_normal(x.shape) # 路径 2:闭式解(用同一步里现造的噪声,保证逐样本可比) xt_cf, _ = q_sample(x0, t_idx, alpha_bar, rng) # 逐样本比不了(噪声不同),比统计量 out.append({ "t": t_idx + 1, "alpha_bar": a, "mean_step": np.abs(x.mean(0) - np.sqrt(a) * x0.mean(0)).max(), "var_step": np.abs(x.var(0) - (a * x0.var(0) + (1 - a))).max(), "var_cf": np.abs(xt_cf.var(0) - (a * x0.var(0) + (1 - a))).max(), }) return out # ────────────────────────────────────────────────────────────────────── # B. 调度对比 # ────────────────────────────────────────────────────────────────────── def schedule_report(): lin_b = linear_beta(T) lin_a = alpha_bar_from_beta(lin_b) cos_b = cosine_beta(T) cos_a = alpha_bar_from_beta(cos_b) rows = [] for t in [1, 50, 100, 250, 500, 750, 900, 1000]: rows.append({ "t": t, "lin_a": lin_a[t - 1], "cos_a": cos_a[t - 1], "lin_snr": snr_db(lin_a[t - 1]), "cos_snr": snr_db(cos_a[t - 1]), "lin_b": lin_b[t - 1], "cos_b": cos_b[t - 1], }) return lin_b, lin_a, cos_b, cos_a, rows def half_life(alpha_bar): """alpha_bar 掉到 0.5 是第几步——一半的采样步数花在这之后。""" idx = np.argmax(alpha_bar < 0.5) return int(idx) + 1 if alpha_bar[idx] < 0.5 else T # ────────────────────────────────────────────────────────────────────── # C. 区分边缘终端 KL 与 ELBO 的条件终端 KL # ────────────────────────────────────────────────────────────────────── def terminal_kl(alpha_bar_T): """KL( q(x_T) || N(0,I) ),q(x_T) 是七分量高斯混合。 没有闭式解,用蒙特卡洛:E_{q(x_T)}[ log q(x_T) - log N(0,I) ] """ rng = np.random.default_rng(SEED + 7) x = sample_data(400_000, rng) w, m, v = noised_mixture(alpha_bar_T) eps = rng.standard_normal(x.shape) xt = np.sqrt(alpha_bar_T) * x + np.sqrt(1.0 - alpha_bar_T) * eps lq = log_density(xt, w, m, v) ln = -0.5 * ((xt ** 2).sum(1) + D * np.log(2 * np.pi)) return float((lq - ln).mean()), float((lq - ln).std() / np.sqrt(len(xt))) def terminal_elbo_kl(alpha_bar_T): """E_data KL(q(x_T|x_0) || N(0,I)) 的闭式值。""" if not 0 <= alpha_bar_T < 1: raise ValueError("alpha_bar_T must lie in [0,1)") second_moment = np.sum(WEIGHTS * (np.sum(MEANS ** 2, axis=1) + D * SCALES ** 2)) return 0.5 * (alpha_bar_T * second_moment - D * alpha_bar_T - D * np.log1p(-alpha_bar_T)) def main(): rng = np.random.default_rng(SEED) print("=" * 68) print("A. 闭式解核对:逐步迭代 vs 一步闭式(20 万样本,2D)") print("=" * 68) lin_b, lin_a, cos_b, cos_a, _rows = schedule_report() for r in check_closed_form(rng, lin_a): print(f" t={r['t']:>4} alpha_bar={r['alpha_bar']:.6e} " f"|均值差|={r['mean_step']:.2e} 逐步方差误={r['var_step']:.2e} " f"闭式方差误={r['var_cf']:.2e}") print() print("=" * 68) print("B. 两种调度把信息怎么抹掉的") print("=" * 68) print(f" {'t':>5} {'alpha_bar(linear)':>18} {'alpha_bar(cosine)':>18} " f"{'SNR_dB(lin)':>12} {'SNR_dB(cos)':>12}") for r in _rows: print(f" {r['t']:>5} {r['lin_a']:>18.6e} {r['cos_a']:>18.6e} " f"{r['lin_snr']:>12.2f} {r['cos_snr']:>12.2f}") print(f"\n alpha_bar 掉到 0.5 的步数:linear = {half_life(lin_a)}," f"cosine = {half_life(cos_a)}") print(f" 也就是说 linear 调度下,{100 * (1 - half_life(lin_a) / T):.1f}% 的步数" f"花在信噪比已经低于 0 dB 的区域") print() print("=" * 68) print("C. 边缘终端 KL 和 ELBO 条件终端 KL(不同对象)") print("=" * 68) for name, aT in [("linear", lin_a[-1]), ("cosine", cos_a[-1])]: kl, se = terminal_kl(aT) print(f" {name:>7}: E_data KL(q(x_T|x_0)||N(0,I)) = {terminal_elbo_kl(aT):.6e} nats") print(f" {name:>7}: alpha_bar_T={aT:.6e} KL(q(x_T)||N(0,I)) = {kl:.6e} nats" f" (±{se:.1e})") print() print("=" * 68) print("D. 连续时间对账:beta(t) = T * beta_i,积分出来的 alpha_bar 对不对") print("=" * 68) b_min, b_max = T * float(lin_b[0]), T * float(lin_b[-1]) print(f" beta_min = {b_min:.4f}(Song et al. VP-SDE 默认 0.1)" f" beta_max = {b_max:.4f}(默认 20)") for t in [0.25, 0.5, 0.75, 1.0]: i = int(round(t * T)) - 1 a_cont = np.exp(-(b_min * t + 0.5 * (b_max - b_min) * t * t)) print(f" t={t:.2f} 离散 alpha_bar={lin_a[i]:.6e} " f"连续 exp(-int beta)={a_cont:.6e} 相对差={abs(lin_a[i] - a_cont) / lin_a[i]:.2e}") print() print("=" * 68) print("E. 终点还剩多少原始信号(非零终端 SNR 问题)") print("=" * 68) for name, aT in [("linear", lin_a[-1]), ("cosine", cos_a[-1])]: amp = np.sqrt(aT) snr_T = aT / (1.0 - aT) print(f" {name:>7}: alpha_bar_T={aT:.6e} sqrt(alpha_bar_T)={amp:.6e}" f" => x_T 里原始信号的振幅占比 {amp * 100:.4f}%") print(f" 终端 SNR = {snr_T:.4e}(= {snr_db(aT):.2f} dB)," f"理论值应当是 0") if __name__ == "__main__": main() score_bridges.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """把 score / epsilon / x_0 / v 四种参数化钉死在同一组恒等式上。 扩散模型的代码里同一个东西有四种写法,换参数化是家常便饭(SD 2.x 用 v, SD 1.x 和大部分视频模型用 epsilon,蒸馏论文里又爱用 x_0)。这篇要说的是: 它们描述的是同一个量,但**作为训练目标并不等价**——换参数化等于给每个 时间步偷偷换了一个权重。 跑法:python score_bridges.py """ import numpy as np from forward_diffusion import ( D, SEED, T, MEANS, SCALES, WEIGHTS, alpha_bar_from_beta, linear_beta, noised_mixture, sample_data, score_fn, ) ALPHA_BAR = alpha_bar_from_beta(linear_beta(T)) def abar(t): """t 是 1-based 步号。""" return 1.0 if t == 0 else float(ALPHA_BAR[t - 1]) def score_at(x, t): return score_fn(x, abar(t)) # ────────────────────────────────────────────────────────────────────── # A. 精确后验均值 E[x_0 | x_t] # ────────────────────────────────────────────────────────────────────── def posterior_mean_exact(x, t): """q(x_0 | x_t) 仍是高斯混合,第 k 个分量: 后验权重 r_k = 与 score 里用的是同一份责任度 后验方差 C_k = (1/s_k^2 + alpha_bar/(1-alpha_bar))^{-1} 后验均值 m_k = C_k * (mu_k / s_k^2 + sqrt(alpha_bar) * x_t / (1-alpha_bar)) """ a = abar(t) if a >= 1.0: return x.copy() _, m_t, v_t = noised_mixture(a) # 加噪后各分量的均值 / 方差 n = x.shape[0] sq = ((x[:, None, :] - m_t[None, :, :]) ** 2).sum(-1) logc = -0.5 * (sq / v_t[None, :] + D * np.log(2 * np.pi * v_t)[None, :]) mx = logc.max(1, keepdims=True) resp = WEIGHTS[None, :] * np.exp(logc - mx) resp = resp / resp.sum(1, keepdims=True) # [n, K] s2 = SCALES ** 2 C = 1.0 / (1.0 / s2 + a / (1.0 - a)) # [K] mk = C[None, :, None] * (MEANS[None, :, :] / s2[None, :, None] + np.sqrt(a) * x[:, None, :] / (1.0 - a)) return (resp[:, :, None] * mk).sum(1) def tweedie_mean(x, t): """Tweedie:E[x_0 | x_t] = (x_t + (1 - alpha_bar_t) * score) / sqrt(alpha_bar_t)。""" a = abar(t) if a >= 1.0: return x.copy() return (x + (1.0 - a) * score_at(x, t)) / np.sqrt(a) def eps_from_score(x, t): """eps = -sqrt(1 - alpha_bar_t) * score。""" a = abar(t) return -np.sqrt(1.0 - a) * score_at(x, t) def eps_from_x0(x, x0_hat, t): """反解:eps = (x_t - sqrt(alpha_bar_t) * x_0) / sqrt(1 - alpha_bar_t)。""" a = abar(t) return (x - np.sqrt(a) * x0_hat) / np.sqrt(1.0 - a) def v_target(x0, eps, t): """v = sqrt(alpha_bar) * eps - sqrt(1 - alpha_bar) * x_0。""" a = abar(t) return np.sqrt(a) * eps - np.sqrt(1.0 - a) * x0 def recover_from_v(x, v, t): """v 和 x_t 之间是一个旋转:x_0 = a*x_t - b*v,eps = b*x_t + a*v。 变换矩阵 [[a, b], [-b, a]] 行列式为 a^2+b^2=1,所以它是正交的—— 这也意味着 v 参数化不会放大噪声。 """ a, b = np.sqrt(abar(t)), np.sqrt(1.0 - abar(t)) return a * x - b * v, b * x + a * v # ────────────────────────────────────────────────────────────────────── # B. 蒙特卡洛交叉验证(不靠闭式解,纯重要性采样) # ────────────────────────────────────────────────────────────────────── def is_posterior_mean(xi, t, n=1_000_000, seed=SEED + 11): """E[x_0 | x_t = xi] 的自归一化重要性采样估计。 x_0^i 就是从 p_0 里采的,所以权重直接取 q(xi | x_0^i) 即可, 不需要知道归一化常数。 """ rng = np.random.default_rng(seed) x0 = sample_data(n, rng) a = abar(t) b2 = 1.0 - a sq = ((xi[None, :] - np.sqrt(a) * x0) ** 2).sum(1) logw = -0.5 * sq / b2 logw -= logw.max() w = np.exp(logw) ess = w.sum() ** 2 / (w ** 2).sum() return (w[:, None] * x0).sum(0) / w.sum(), ess # ────────────────────────────────────────────────────────────────────── # C. 三种参数化在 epsilon 空间下的每步权重 # ────────────────────────────────────────────────────────────────────── def loss_weights(t): """把三种参数化的训练损失都换算回"相当于给 eps 误差加了多大权重"。 - eps 参数化:delta_eps = delta_u -> 权重 1 - x_0 参数化:eps = (x_t - sqrt(a) x_0)/sqrt(1-a) delta_eps = -sqrt(a/(1-a)) * delta_u -> 权重 a/(1-a) = SNR 反过来,x_0 的损失 ||delta_u||^2 折算成 eps 误差是 ||delta_eps||^2 = SNR * ||delta_x0||^2, 即 x_0 损失对 eps 误差的权重是 1/SNR - v 参数化 :eps = b*x_t + a*v,delta_eps = a * delta_v -> v 损失折算成 eps 误差的权重是 1/a = 1/alpha_bar 返回 (权重_eps, 权重_v, 权重_x0),都是"乘在 ||delta_eps||^2 上的系数"。 """ a = abar(t) if t > 0 else 1.0 snr = a / (1.0 - a) return 1.0, 1.0 / a, 1.0 / snr def main(): rng = np.random.default_rng(SEED) print("=" * 72) print("A. Tweedie 恒等式核对:闭式后验均值 vs 公式 (x_t + (1-a) s)/sqrt(a)") print("=" * 72) print(f" {'t':>5} {'alpha_bar':>14} {'最大绝对误差':>16} {'x_t 范数均值':>14}") for t in [1, 50, 100, 250, 500, 750, 1000]: x = sample_data(2000, rng) xt, _ = (np.sqrt(abar(t)) * x + np.sqrt(1 - abar(t)) * rng.standard_normal(x.shape), None) e1 = posterior_mean_exact(xt, t) e2 = tweedie_mean(xt, t) err = np.abs(e1 - e2).max() print(f" {t:>5} {abar(t):>14.6e} {err:>16.3e} " f"{np.linalg.norm(xt, axis=1).mean():>14.4f}") print() print("=" * 72) print("B. eps 与 score 的换算:eps = -sqrt(1-a) * s,两条路算 eps 对不对") print("=" * 72) print(f" {'t':>5} {'|eps(score) - eps(x_0)| 最大':>28}") for t in [50, 250, 500, 1000]: x0 = sample_data(2000, rng) eps_true = rng.standard_normal(x0.shape) a = abar(t) xt = np.sqrt(a) * x0 + np.sqrt(1 - a) * eps_true # 路线 1:从 score e_s = eps_from_score(xt, t) # 路线 2:从 Tweedie 反解出的 x_0 e_x = eps_from_x0(xt, tweedie_mean(xt, t), t) print(f" {t:>5} {np.abs(e_s - e_x).max():>28.3e}") # 顺便看看 MMSE 估计量离真实 eps 有多远(这是"复原不了"的那部分) print(f" (MMSE 估计量与本次真实 eps 的 RMSE = " f"{np.sqrt(((e_s - eps_true) ** 2).sum(1).mean()):.4f}," f"sqrt(2D)={np.sqrt(2 * D):.4f} 是纯瞎猜的水平)") print() print("=" * 72) print("C. v 参数化:它与 (x_t, eps, x_0) 之间是一个旋转") print("=" * 72) for t in [50, 500, 1000]: x0 = sample_data(2000, rng) eps = rng.standard_normal(x0.shape) a = abar(t) xt = np.sqrt(a) * x0 + np.sqrt(1 - a) * eps v = v_target(x0, eps, t) x0_rec, eps_rec = recover_from_v(xt, v, t) print(f" t={t:>5} 重建 x_0 误差={np.abs(x0_rec - x0).max():.3e} " f"重建 eps 误差={np.abs(eps_rec - eps).max():.3e} " f"v 的范数均值={np.linalg.norm(v, axis=1).mean():.4f}") print() print("=" * 72) print("D. 蒙特卡洛交叉验证(100 万样本的重要性采样,不依赖上面的闭式解)") print("=" * 72) print(f" {'t':>5} {'Tweedie':>22} {'重要性采样':>22} {'有效样本数':>12}") for t in [50, 250, 500, 1000]: xi = sample_data(1, rng)[0] xi = np.sqrt(abar(t)) * xi + np.sqrt(1 - abar(t)) * rng.standard_normal(2) tm = tweedie_mean(xi[None, :], t)[0] im, ess = is_posterior_mean(xi, t) print(f" {t:>5} {np.array2string(tm, precision=4):>22} " f"{np.array2string(im, precision=4):>22} {ess:>12.0f}") print() print("=" * 72) print("E. 换参数化 = 给每个时间步换权重(折算回 eps 空间的系数)") print("=" * 72) print(f" {'t':>5} {'alpha_bar':>13} {'w_eps':>12} {'w_v':>14} {'w_x0':>14}") for t in [1, 10, 50, 100, 250, 500, 750, 900, 1000]: we, wv, wx = loss_weights(t) print(f" {t:>5} {abar(t):>13.6e} {we:>12.4f} {wv:>14.4e} {wx:>14.4e}") wes, wvs, wxs = zip(*[loss_weights(t) for t in range(1, T + 1)]) print(f"\n eps 参数化: 权重恒为 1,动态范围 {max(wes) / min(wes):.1f}") print(f" v 参数化: 权重 {min(wvs):.4e} ~ {max(wvs):.4e}," f"动态范围 {max(wvs) / min(wvs):.3e}") print(f" x_0 参数化: 权重 {min(wxs):.4e} ~ {max(wxs):.4e}," f"动态范围 {max(wxs) / min(wxs):.3e}") if __name__ == "__main__": main() make_figures.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """生成「扩散过程的前向与反向推导」的四张解释图。 数值一律从同目录的三个脚本里取(forward_diffusion / reverse_sampling / score_bridges),这里只负责画——改了那边这里要重跑,免得图和正文数字打架。 四张图分别回答: 1. 噪声调度到底把"难度"怎么分配到 1000 步上的 2. 五种采样器的误差随步数怎么降,以及每步跨过的积分量有多大 3. 前向把数据流推成球、反向沿轨迹走回七个模式,长什么样 4. 换参数化等价于给每个时间步换了多大权重(跨 10 个数量级) 只依赖 numpy + matplotlib。跑法:python make_figures.py """ import textwrap from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np from matplotlib.collections import LineCollection import forward_diffusion as FD import reverse_sampling as RS import score_bridges as SB ROOT = Path(__file__).resolve().parents[1] OUT = ROOT / "figures" OUT.mkdir(parents=True, exist_ok=True) plt.rcParams.update({ "font.sans-serif": ["PingFang SC", "Arial Unicode MS", "DejaVu Sans"], "axes.unicode_minus": False, "figure.dpi": 160, "savefig.bbox": "tight", }) INK = "#1f2937" MUTE = "#6b7280" C_LIN = "#d1495b" # 线性调度 / 差的一侧:红 C_COS = "#2f6fb0" # 余弦调度:蓝 C_DDPM = "#e0a03c" C_DDIM = "#2f9e6f" C_EM = "#8b5cf6" C_ODE = "#d1495b" C_HEUN = "#1f6feb" C_OK = "#2f9e6f" C_BAD = "#d1495b" def sci(v): """2.48e+08 -> 2.5e8,读起来省事。""" return f"{v:.1e}".replace("e+0", "e").replace("e+", "e").replace("e-0", "e-") def style(ax, title, xlabel=None, ylabel=None): ax.set_title(title, fontsize=11.5, color=INK, pad=10, loc="left") if xlabel: ax.set_xlabel(xlabel, fontsize=10, color=MUTE) if ylabel: ax.set_ylabel(ylabel, fontsize=10, color=MUTE) ax.tick_params(colors=MUTE, labelsize=9) for s in ("top", "right"): ax.spines[s].set_visible(False) for s in ("left", "bottom"): ax.spines[s].set_color("#d1d5db") ax.grid(alpha=0.25, linewidth=0.6) ax.set_axisbelow(True) def footer(fig, text, width=118): """把「这张图要看什么」放到坐标轴下方。 必须放在 y<0 的位置:bbox_inches="tight" 会把负坐标的 artist 一起收进来, 放在 0~0.05 之间的话会和 x 轴标签叠在一起。 """ wrapped = "\n".join(textwrap.wrap(text, width=width)) fig.text(0.012, -0.13, wrapped, fontsize=8.5, color=MUTE, va="top", ha="left", linespacing=1.6) # ────────────────────────────────────────────────────────────────────── # 图 1:噪声调度怎么分配难度 # ────────────────────────────────────────────────────────────────────── def fig_schedule(): lin_b, lin_a, cos_b, cos_a, rows = FD.schedule_report() ts = np.arange(1, FD.T + 1) snr_lin = FD.snr_db(lin_a) snr_cos = FD.snr_db(cos_a) fig, axes = plt.subplots(1, 2, figsize=(11.2, 4.1)) ax = axes[0] ax.plot(ts, lin_a, color=C_LIN, lw=2.0, label=r"linear $\beta$ 调度") ax.plot(ts, cos_a, color=C_COS, lw=2.0, label="余弦调度") ax.set_yscale("log") ax.axhline(0.5, color=MUTE, ls=":", lw=1.0) for x, c, lbl in [(FD.half_life(lin_a), C_LIN, f"linear 跌破 0.5:第 {FD.half_life(lin_a)} 步"), (FD.half_life(cos_a), C_COS, f"余弦跌破 0.5:第 {FD.half_life(cos_a)} 步")]: ax.axvline(x, color=c, ls="--", lw=1.0, alpha=0.8) ax.legend(fontsize=8.5, frameon=False, loc="lower left") style(ax, r"$\bar\alpha_t$:还剩下多少原始信号", "步数 t", r"$\bar\alpha_t$(对数轴)") ax = axes[1] ax.plot(ts, snr_lin, color=C_LIN, lw=2.0, label="linear") ax.plot(ts, snr_cos, color=C_COS, lw=2.0, label="余弦") ax.axhline(0.0, color=INK, lw=1.0) ax.fill_between(ts, snr_lin.min(), 0.0, where=(snr_lin < 0), color=C_LIN, alpha=0.10) i500 = 499 ax.annotate(f"t=500\nlinear {snr_lin[i500]:.1f} dB\n余弦 {snr_cos[i500]:.1f} dB", xy=(500, snr_lin[i500]), xytext=(620, -6), fontsize=8.5, color=INK, arrowprops=dict(arrowstyle="->", color=MUTE, lw=0.9)) ax.legend(fontsize=8.5, frameon=False, loc="upper right") style(ax, "信噪比:低于 0 dB 时信号已经被噪声淹没", "步数 t", "SNR (dB)") fig.suptitle("图 1 线性调度把 74% 的步数花在了信噪比已经低于 0 dB 的区域", fontsize=12, color=INK, x=0.012, ha="left", y=1.02) footer(fig, "要看什么:左图是「还剩多少信号」,右图是「信噪比」。线性调度在第 260 步就把一半信号丢光了," "后续处在较低 SNR 区间,但不等于无效计算;图中 SNR 以数据方差为 1 作参考。") fig.savefig(OUT / "schedule_snr.png") plt.close(fig) # ────────────────────────────────────────────────────────────────────── # 图 2:采样器误差随步数怎么降 + 每步跨过的积分量 # ────────────────────────────────────────────────────────────────────── def fig_samplers(res=None): if res is None: res, _ = RS.sweep() n_list = [10, 20, 25, 50, 100, 200, 1000] cols = {"ddpm": C_DDPM, "ddim": C_DDIM, "em": C_EM, "ode": C_ODE, "heun": C_HEUN} names = {"ddpm": "DDPM 祖采样", "ddim": "DDIM (η=0)", "em": "Euler–Maruyama (SDE)", "ode": "Euler (概率流 ODE)", "heun": "Heun 二阶 (ODE)"} fig, axes = plt.subplots(1, 2, figsize=(11.2, 4.3)) ax = axes[0] for name in ["ddpm", "ddim", "em", "ode", "heun"]: y = [abs(res[(name, n)]["dlogp"]) for n in n_list] ax.plot(n_list, y, marker="o", ms=4.5, lw=1.9, color=cols[name], label=names[name]) ax.set_xscale("log") ax.set_yscale("log") ax.axhline(0.0528, color=MUTE, ls="--", lw=1.0) ax.text(11, 0.062, "一次真实样本对照差值 0.053", fontsize=8, color=MUTE) ax.legend(fontsize=8, frameon=False, loc="upper right") style(ax, "单一统计量偏差(不能独立确认分布正确)", "采样步数 N(对数轴)", r"$|$mean log $p_0$ 偏移$|$(nats)") ax = axes[1] st = RS.stride_stats() ns = [r["n"] for r in st] ax.plot(ns, [r["L_max"] for r in st], marker="s", ms=4.5, lw=1.9, color=C_ODE, label=r"每步跨过的积分量 $L=\int\beta(u)du$(最大)") ax.plot(ns, [r["euler_err"] * 100 for r in st], marker="^", ms=4.5, lw=1.9, color=C_HEUN, label=r"Euler 近似 $e^{L/2}\approx 1+L/2$ 的相对误差") ax.set_xscale("log") ax.set_yscale("log") for r in st: if r["n"] in (10, 50): ax.annotate(f"{r['euler_err'] * 100:.1f}%", xy=(r["n"], r["euler_err"] * 100), xytext=(r["n"] * 1.15, r["euler_err"] * 100 * 1.6), fontsize=8.5, color=C_HEUN) ax.legend(fontsize=8, frameon=False, loc="upper right") style(ax, "为什么朴素 Euler 在大步长下会崩", "采样步数 N(对数轴)", "数值(对数轴,% 按数值读)") fig.suptitle("图 2 采样器诊断统计量与线性漂移的离散误差", fontsize=12, color=INK, x=0.012, ha="left", y=1.02) footer(fig, "要看什么:左图为单次实验,近零差异不作显著性排名;右图只比较线性漂移项的数值近似。DDPM/DDIM 把线性部分" "精确解掉了(直接乘 √ᾱ),Euler 却把它展开成一阶,N=10 时每步误差就有 25%。") fig.savefig(OUT / "sampler_scaling.png") plt.close(fig) return res # ────────────────────────────────────────────────────────────────────── # 图 3:前向流形被推成球,反向轨迹走回七个模式 # ────────────────────────────────────────────────────────────────────── def fig_trajectories(): rng = np.random.default_rng(FD.SEED) ab = FD.alpha_bar_from_beta(FD.linear_beta(FD.T)) x0 = FD.sample_data(600, rng) snaps = [0, 60, 200, 500, 1000] cols_f = plt.cm.viridis(np.linspace(0.05, 0.85, len(snaps))) # 反向:从同一批噪声出发,用 Heun 走 50 步,记录轨迹 x = rng.standard_normal((40, FD.D)) taus = RS.make_stride(50, FD.T) paths = [x.copy()] for j in range(len(taus) - 1): L = RS._log_step(ab, taus[j], taus[j + 1]) if L < 1e-12: continue s0 = RS.score_at(x, ab, taus[j]) d0 = 0.5 * L * (x + s0) x1 = x + d0 s1 = RS.score_at(x1, ab, taus[j + 1]) x = x + 0.5 * (d0 + 0.5 * L * (x1 + s1)) paths.append(x.copy()) fig, axes = plt.subplots(1, 2, figsize=(11.4, 4.9)) ax = axes[0] for t, c in zip(snaps, cols_f): a = 1.0 if t == 0 else float(ab[t - 1]) xt = np.sqrt(a) * x0 + np.sqrt(1 - a) * rng.standard_normal(x0.shape) ax.scatter(xt[:, 0], xt[:, 1], s=7, color=c, alpha=0.65, label=f"t={t}") ax.scatter(FD.MEANS[:, 0], FD.MEANS[:, 1], marker="x", s=70, color=INK, linewidths=1.6, label="分量中心") ax.legend(fontsize=8, frameon=False, loc="upper left", ncol=2) style(ax, "前向:七个模式被逐步抹成一个标准正态球", r"$x_1$", r"$x_2$") ax.set_aspect("equal") ax = axes[1] # 背景:p_0 的密度等高线 g = np.linspace(-4.6, 4.6, 220) GX, GY = np.meshgrid(g, g) grid = np.stack([GX.ravel(), GY.ravel()], 1) lp = FD.log_density(grid, FD.WEIGHTS, FD.MEANS, FD.SCALES ** 2) dens = np.exp(lp - lp.max()).reshape(GX.shape) ax.contourf(GX, GY, dens, levels=np.linspace(0.02, 1.0, 12), cmap="Blues", alpha=0.9, vmin=0.0, vmax=1.6) ax.contour(GX, GY, dens, levels=[0.05, 0.2, 0.5], colors="#1d4ed8", linewidths=0.7, alpha=0.55) P = np.stack(paths, 1) # [n, steps, 2] for i in range(P.shape[0]): seg = P[i] ax.add_collection(LineCollection( [seg[k:k + 2] for k in range(len(seg) - 1)], colors="#1f6feb", linewidths=0.8, alpha=0.5)) ax.scatter(P[:, 0, 0], P[:, 0, 1], s=14, color=MUTE, label="起点 x_T ~ N(0,I)") ax.scatter(P[:, -1, 0], P[:, -1, 1], s=16, color=C_OK, label="终点 x_0") ax.scatter(FD.MEANS[:, 0], FD.MEANS[:, 1], marker="x", s=70, color=INK, linewidths=1.6) ax.legend(fontsize=8, frameon=False, loc="upper left") style(ax, "反向:50 步 Heun,40 条轨迹从噪声回到模式", r"$x_1$", r"$x_2$") ax.set_aspect("equal") fig.suptitle("图 3 前向是「加水搅匀」,反向是「沿着 score 场把水滤掉」", fontsize=12, color=INK, x=0.012, ha="left", y=1.0) footer(fig, "要看什么:左边 t=200 时七个模式已经互相渗透,t=500 彻底成球;右边每条轨迹在最后十几步才" "「决定」进哪个模式——前面的漫长路程都在把粗轮廓搭起来,这解释了为什么大步长主要伤细节。") fig.savefig(OUT / "trajectories.png") plt.close(fig) # ────────────────────────────────────────────────────────────────────── # 图 4:换参数化 = 换每步权重 # ────────────────────────────────────────────────────────────────────── def fig_param_weights(): ts = np.arange(1, FD.T + 1) w = np.array([SB.loss_weights(t) for t in ts]) # [T, 3] fig, ax = plt.subplots(figsize=(7.6, 4.5)) ax.plot(ts, w[:, 0], color=C_OK, lw=2.4, label=r"$\varepsilon$ 参数化:恒为 1") ax.plot(ts, w[:, 1], color=C_DDPM, lw=2.0, label=r"$v$ 参数化:$1/\bar\alpha_t$") ax.plot(ts, w[:, 2], color=C_BAD, lw=2.0, label=r"$x_0$ 参数化:$1/\mathrm{SNR}_t$") ax.set_yscale("log") ax.axhline(1.0, color=MUTE, ls=":", lw=1.0) ax.annotate(rf"$x_0$ 跨 {sci(w[-1, 2] / w[0, 2])} 倍(1.0e-4 → 2.5e4)", xy=(1000, w[-1, 2]), xytext=(300, 6e2), fontsize=9, color=C_BAD, arrowprops=dict(arrowstyle="->", color=C_BAD, lw=0.9)) ax.annotate(rf"$v$ 跨 {sci(w[-1, 1] / w[0, 1])} 倍(1.0 → 2.5e4)", xy=(1000, w[-1, 1]), xytext=(300, 1.6e0), fontsize=9, color=C_DDPM, arrowprops=dict(arrowstyle="->", color=C_DDPM, lw=0.9)) ax.set_ylim(5e-5, 1e5) ax.legend(fontsize=9, frameon=False, loc="lower right") style(ax, "同一个模型,换参数化等于给每个时间步换权重", "步数 t", r"折算到 $\varepsilon$ 空间后的权重(对数轴)") fig.suptitle(r"图 4 $\varepsilon$ 参数化的权重恒为 1,$x_0$ 参数化跨了 8 个数量级", fontsize=12, color=INK, x=0.012, ha="left", y=1.03) footer(fig, "要看什么:三种参数化描述的是同一个量,但作为训练目标并不等价。$x_0$ 参数化在 t=1 处权重只有" " 1e-4、在 t=1000 处却有 2.5e4。这些是代数残差权重,不是实测参数梯度;不能据此单独判断 " "训练稳定性或最终质量。") fig.savefig(OUT / "param_weights.png") plt.close(fig) def main(): fig_schedule() res = fig_samplers() fig_trajectories() fig_param_weights() print(f"[OK] 四张图已写入 {OUT}") for f in sorted(OUT.glob("*.png")): print(f" {f.name} {f.stat().st_size / 1024:.0f} KB") if __name__ == "__main__": main() reverse_sampling.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """反向采样:五种采样器在"同一个 oracle score"下的正面对比。 这一节要回答的是:既然反向过程的 drift 只有唯一一种写法,为什么工业界 能搞出 DDPM / DDIM / DPM-Solver / Euler / Heun 这么多种采样器?差别到底 在哪里? 关键设计:目标分布用 2D 高斯混合,加噪之后仍然是高斯混合,所以 nabla_x log p_t(x) 有闭式解。也就是说这里的 score 是**理论最优**的, 不掺任何网络拟合误差——差异来自离散化、有限终端噪声近似以及蒙特卡洛波动。 五种采样器: ddpm DDPM 祖采样(ancestral),高斯反向核近似,每步注噪声 ddim DDIM,eta=0,确定性 em Euler-Maruyama 解反向 SDE,一阶,注噪声 ode Euler 解概率流 ODE,一阶,确定性 heun Heun 二阶解概率流 ODE,每步两次 score 评估 运行: python reverse_sampling.py """ import numpy as np from forward_diffusion import ( D, SEED, T, MEANS, SCALES, WEIGHTS, alpha_bar_from_beta, cosine_beta, linear_beta, log_density, sample_data, score_fn, ) # ────────────────────────────────────────────────────────────────────── # 时间步工具 # ────────────────────────────────────────────────────────────────────── def make_stride(n_steps, T=T): """把 0..T 均匀切成 n_steps 段,返回 1-based 的 t 序列(含 T 与 0)。 例:n_steps=50, T=1000 -> [1000, 980, 960, ..., 20, 0] """ idx = np.linspace(T, 0, n_steps + 1).astype(int) # 保证严格递减且唯一 idx = np.unique(idx)[::-1] return list(idx) def abar_at(alpha_bar, t): """t 是 1-based 步号;t=0 时 alpha_bar=1(即 x_0 本身)。""" return 1.0 if t == 0 else float(alpha_bar[t - 1]) def score_at(x, alpha_bar, t): return score_fn(x, abar_at(alpha_bar, t)) # ────────────────────────────────────────────────────────────────────── # 五种采样器 # ────────────────────────────────────────────────────────────────────── def _tweedie_x0(x, s, a_cur): """Tweedie 公式:E[x_0 | x_t] = (x_t + (1-a_t) * score) / sqrt(a_t)。""" return (x + (1.0 - a_cur) * s) / np.sqrt(a_cur) def _eps_from_score(s, a_cur): """epsilon 与 score 的换算:eps = -sqrt(1 - a_t) * score。""" return -np.sqrt(1.0 - a_cur) * s def run_ddpm(x, alpha_bar, taus, rng, **kw): """DDPM 祖采样:用精确反向均值配固定高斯方差,有限步并非精确逆核。 x_{t-1} = (x_t + beta_t * s) / sqrt(alpha_t) + sigma_t * z sigma_t^2 = beta_t * (1 - abar_{t-1}) / (1 - abar_t) """ for j in range(len(taus) - 1): t_cur, t_prev = taus[j], taus[j + 1] a_cur, a_prev = abar_at(alpha_bar, t_cur), abar_at(alpha_bar, t_prev) alpha_j = a_cur / a_prev beta_j = 1.0 - alpha_j if beta_j < 1e-12: continue s = score_at(x, alpha_bar, t_cur) mean = (x + beta_j * s) / np.sqrt(alpha_j) var = beta_j * (1.0 - a_prev) / (1.0 - a_cur) x = mean + np.sqrt(max(var, 0.0)) * rng.standard_normal(x.shape) return x def run_ddim(x, alpha_bar, taus, rng, **kw): """DDIM,eta=0:先跳到去噪后的 x_0,再重新加回 target 时刻的噪声。 x_{t-1} = sqrt(abar_{t-1}) * xhat_0 + sqrt(1 - abar_{t-1}) * eps_hat """ for j in range(len(taus) - 1): t_cur, t_prev = taus[j], taus[j + 1] a_cur, a_prev = abar_at(alpha_bar, t_cur), abar_at(alpha_bar, t_prev) s = score_at(x, alpha_bar, t_cur) xhat0 = _tweedie_x0(x, s, a_cur) eps_hat = _eps_from_score(s, a_cur) x = np.sqrt(a_prev) * xhat0 + np.sqrt(max(1.0 - a_prev, 0.0)) * eps_hat return x def _log_step(alpha_bar, t_cur, t_prev): """跨一步"累积"起来的积分量 L = int beta du = log(alpha_bar_{prev}/alpha_bar_{cur}) > 0。 这一步很关键:正向 SDE 的时间是 u = t/T,从 0 涨到 1;反向采样是让 u **往回走**,所以 du 是负的。把 du = -1/N 代进去之后,drift 的符号会整体 翻过来——反向过程是把被压扁的分布"吹"回原样,而不是继续压。 用 L 记这段区间上 beta 的积分,Euler 步就写成: dx = +L * (0.5 x + c * s) (c=1 是 SDE,c=1/2 是 ODE) """ a_cur, a_prev = abar_at(alpha_bar, t_cur), abar_at(alpha_bar, t_prev) return float(np.log(a_prev / a_cur)) def run_em(x, alpha_bar, taus, rng, **kw): """Euler-Maruyama 解反向 SDE:全量 score + 注入噪声。 反向 SDE 的 drift 里 score 的系数是 **1 倍**(不是 ODE 的半倍), 另外每步还要加 sqrt(L) 的噪声。 """ for j in range(len(taus) - 1): L = _log_step(alpha_bar, taus[j], taus[j + 1]) if L < 1e-12: continue s = score_at(x, alpha_bar, taus[j]) x = x + L * (0.5 * x + s) + np.sqrt(L) * rng.standard_normal(x.shape) return x def run_ode(x, alpha_bar, taus, rng, **kw): """Euler 解概率流 ODE:半倍 score,无噪声。 和 EM 只差两处:score 系数从 1 减到 0.5,噪声项整个去掉。 """ for j in range(len(taus) - 1): L = _log_step(alpha_bar, taus[j], taus[j + 1]) if L < 1e-12: continue s = score_at(x, alpha_bar, taus[j]) x = x + 0.5 * L * (x + s) return x def run_heun(x, alpha_bar, taus, rng, **kw): """Heun 二阶解概率流 ODE:Euler 预测一步,再用终点的 score 校正一次。 每步两次 score 评估(NFE = 2 * n_steps)。 """ for j in range(len(taus) - 1): L = _log_step(alpha_bar, taus[j], taus[j + 1]) if L < 1e-12: continue s0 = score_at(x, alpha_bar, taus[j]) d0 = 0.5 * L * (x + s0) x1 = x + d0 # Euler 预测 s1 = score_at(x1, alpha_bar, taus[j + 1]) # 在目标时刻再评估一次 d1 = 0.5 * L * (x1 + s1) x = x + 0.5 * (d0 + d1) return x SAMPLERS = { "ddpm": (run_ddpm, 1), "ddim": (run_ddim, 1), "em": (run_em, 1), "ode": (run_ode, 1), "heun": (run_heun, 2), } # ────────────────────────────────────────────────────────────────────── # 评价:MMD^2(RBF 核,median heuristic) # ────────────────────────────────────────────────────────────────────── def _rbf_kernels(a, b, bw): sa = (a ** 2).sum(1) sb = (b ** 2).sum(1) d2 = sa[:, None] + sb[None, :] - 2.0 * a @ b.T return np.exp(-d2 / (2.0 * bw ** 2)) def mmd2(x, y, bw=None, n_cap=1500): """MMD^2 的无偏估计。x 是生成样本,y 是真实样本。""" rng = np.random.default_rng(0) if len(x) > n_cap: x = x[rng.choice(len(x), n_cap, replace=False)] if len(y) > n_cap: y = y[rng.choice(len(y), n_cap, replace=False)] if bw is None: allp = np.vstack([x[:800], y[:800]]) d2 = ((allp[:, None, :] - allp[None, :, :]) ** 2).sum(-1) bw = float(np.sqrt(np.median(d2) / 2.0)) bw = max(bw, 1e-3) kxx = _rbf_kernels(x, x, bw) kyy = _rbf_kernels(y, y, bw) kxy = _rbf_kernels(x, y, bw) n, m = len(x), len(y) # 去掉对角线才无偏 t1 = (kxx.sum() - np.trace(kxx)) / (n * (n - 1)) t2 = (kyy.sum() - np.trace(kyy)) / (m * (m - 1)) t3 = kxy.mean() return float(t1 + t2 - 2 * t3), bw def mean_logp(x): """生成样本在真实 p_0 下的平均 log 密度。 比 MMD 更能抓一种特定的失败:采样器把样本堆到高密度区(过聚拢), 或者样本飘到分量之间的低密度地带。 """ return float(log_density(x, WEIGHTS, MEANS, SCALES ** 2).mean()) def assign_components(x): """按最近的均值把样本归到混合分量上(分量标准差同量级,够用)。""" d = ((x[:, None, :] - MEANS[None, :, :]) ** 2).sum(-1) return d.argmin(1) def tv_weights(x): """经验分量占比与真实权重之间的全变差距离。 抓的是"各模式的比例对不对"——模式丢了一个、或者某个模式被过度采样, 这个数会涨,而 mean log p 未必涨(甚至可能更漂亮)。 """ idx = assign_components(x) n = len(x) hist = np.bincount(idx, minlength=len(WEIGHTS)) / n return float(0.5 * np.abs(hist - WEIGHTS).sum()) def mode_spread_err(x): """每个分量内部样本的均方半径,与理论值 D*s_k^2 的相对误差(取各分量最大)。 抓的是"样本落在模式的中心但挤成一团"或者"散得太开"。 """ idx = assign_components(x) errs = [] for k in range(len(WEIGHTS)): sel = x[idx == k] if len(sel) < 30: errs.append(1.0) continue r2 = ((sel - MEANS[k]) ** 2).sum(1).mean() errs.append(abs(r2 / (D * SCALES[k] ** 2) - 1.0)) return float(max(errs)) def moment_err(x, y): """均值与协方差的绝对误差(作为 MMD 之外的直观补充)。""" return float(np.abs(x.mean(0) - y.mean(0)).max()), \ float(np.abs(np.cov(x.T) - np.cov(y.T)).max()) # ────────────────────────────────────────────────────────────────────── # 主实验 # ────────────────────────────────────────────────────────────────────── def sweep(n_list=(10, 20, 25, 50, 100, 200, 1000), n_samples=4000, schedule="linear", seed=SEED): """固定一份真实样本做参照,扫采样器 × 步数。""" alpha_bar = (alpha_bar_from_beta(linear_beta(T)) if schedule == "linear" else alpha_bar_from_beta(cosine_beta(T))) ref = sample_data(8000, np.random.default_rng(seed)) # 噪声地板:两份**独立**真实样本之间的 MMD^2。 # 单次值不是置信区间;细微差异需要重复抽样评估。 ref_b = sample_data(4000, np.random.default_rng(seed + 999)) floor, bw = mmd2(ref[:4000], ref_b) # 三个主指标都先在大样本真实数据上算一遍当基线,后面一律报"相对基线的偏移", # 理想期望下匹配的统计量差为0,有限样本仍有波动,0也不保证整个分布匹配。 base = { "logp": mean_logp(ref), "tv": tv_weights(ref), "spread": mode_spread_err(ref), } # 一次对照差值(不是标准误或置信区间) noise = { "logp": mean_logp(ref_b) - base["logp"], "tv": tv_weights(ref_b) - base["tv"], "spread": mode_spread_err(ref_b) - base["spread"], } results = {"_floor": floor, "_bw": bw, "_base": base, "_noise": noise} for name, (fn, mult) in SAMPLERS.items(): for n in n_list: rng = np.random.default_rng(seed + 1) x = rng.standard_normal((n_samples, D)) taus = make_stride(n, T) x = fn(x, alpha_bar, taus, rng) m2, _ = mmd2(x, ref, bw=bw) me, ce = moment_err(x, ref) results[(name, n)] = { "mmd2": m2, "nfe": n * mult, "mean_err": me, "cov_err": ce, "dlogp": mean_logp(x) - base["logp"], "dtv": tv_weights(x) - base["tv"], "dspread": mode_spread_err(x) - base["spread"], "finite": bool(np.isfinite(x).all()), } return results, bw def stride_stats(n_list=(10, 20, 50, 100, 200, 1000), T=T): """每一步跨过的"积分量" L = int beta du,以及 Euler 近似 e^{L/2} 会错多少。 DDPM / DDIM 的更新里线性部分是**精确**解掉的(直接乘 sqrt(alpha_bar)), Euler 却把它展开成一阶:e^{L/2} ≈ 1 + L/2。步长越大这两个差得越远, 这就是 N=10 时 Euler 系采样器全线崩掉、DDIM 却还能看的根本原因。 """ alpha_bar = alpha_bar_from_beta(linear_beta(T)) out = [] for n in n_list: taus = make_stride(n, T) Ls = np.array([_log_step(alpha_bar, taus[j], taus[j + 1]) for j in range(len(taus) - 1)]) lm = float(Ls.max()) err = float(np.abs(np.exp(lm / 2) - (1 + lm / 2)) / np.exp(lm / 2)) out.append({"n": n, "L_max": lm, "L_mean": float(Ls.mean()), "euler_err": err}) return out def ablation_noise(n=50, n_samples=4000, seed=SEED): """把 DDIM 的 eta 从 0 拉到 1,看"注入噪声"这一项单独值多少钱。 eta=0 -> DDIM(确定性);eta=1 -> 等价 DDPM 祖采样。 sigma_t = eta * sqrt(beta_t * (1 - abar_{t-1}) / (1 - abar_t)) """ alpha_bar = alpha_bar_from_beta(linear_beta(T)) ref = sample_data(8000, np.random.default_rng(seed)) base_logp, base_tv, bw0 = _ablation_baseline(seed) out = [] taus = make_stride(n, T) for eta in [0.0, 0.25, 0.5, 0.75, 1.0]: rng = np.random.default_rng(seed + 1) x = rng.standard_normal((n_samples, D)) for j in range(len(taus) - 1): t_cur, t_prev = taus[j], taus[j + 1] a_cur, a_prev = abar_at(alpha_bar, t_cur), abar_at(alpha_bar, t_prev) alpha_j = a_cur / a_prev beta_j = 1.0 - alpha_j if beta_j < 1e-12: continue s = score_at(x, alpha_bar, t_cur) xhat0 = _tweedie_x0(x, s, a_cur) eps_hat = _eps_from_score(s, a_cur) var = beta_j * (1.0 - a_prev) / (1.0 - a_cur) sig = eta * np.sqrt(max(var, 0.0)) coeff = np.sqrt(max(1.0 - a_prev - sig ** 2, 0.0)) x = np.sqrt(a_prev) * xhat0 + coeff * eps_hat + sig * rng.standard_normal(x.shape) m2, _ = mmd2(x, ref, bw=bw0) out.append({"eta": eta, "mmd2": m2, "dlogp": mean_logp(x) - base_logp, "dtv": tv_weights(x) - base_tv}) return out def _ablation_baseline(seed=SEED): ref = sample_data(8000, np.random.default_rng(seed)) _, bw0 = mmd2(ref[:2000], sample_data(2000, np.random.default_rng(seed + 3))) return mean_logp(ref), tv_weights(ref), bw0 def main(): print("=" * 74) print("主实验:同一个 oracle score,五种采样器 × 七档步数") print(f"目标:2D 七分量高斯混合;参照集 8000 真实样本;每档生成 {4000} 个") print("主指标是 dlogp / dTV / dspread 三个(都扣掉了真实样本自己的基线," "0 = 完美)") print("MMD^2 只列出来做旁证——它自身的噪声有 ±3.6e-4,见 mmd_noise_check.py") print("=" * 74) res, bw = sweep() floor = res["_floor"] base, noise = res["_base"], res["_noise"] n_list = [10, 20, 25, 50, 100, 200, 1000] print(f" MMD^2(噪声 ±3.6e-4,所以微小差异需重复抽样确认):") print(f" {'采样器':<8}" + "".join(f"{'N=' + str(n):>13}" for n in n_list)) for name in ["ddpm", "ddim", "em", "ode", "heun"]: row = f" {name:<8}" for n in n_list: row += f"{res[(name, n)]['mmd2']:>13.2e}" print(row) print() print(" 主指标 1 —— mean log p_0 相对真实样本的偏移(nats,0 = 完美):") print(f" {'采样器':<8}" + "".join(f"{'N=' + str(n):>13}" for n in n_list)) for name in ["ddpm", "ddim", "em", "ode", "heun"]: row = f" {name:<8}" for n in n_list: row += f"{res[(name, n)]['dlogp']:>13.3f}" print(row) print() print(" 主指标 2 —— 分量占比的全变差距离偏移(0 = 各模式比例都对):") print(f" {'采样器':<8}" + "".join(f"{'N=' + str(n):>13}" for n in n_list)) for name in ["ddpm", "ddim", "em", "ode", "heun"]: row = f" {name:<8}" for n in n_list: row += f"{res[(name, n)]['dtv']:>13.4f}" print(row) print() print(" 主指标 3 —— 模式内均方半径的相对误差(0 = 胖瘦都正好):") print(f" {'采样器':<8}" + "".join(f"{'N=' + str(n):>13}" for n in n_list)) for name in ["ddpm", "ddim", "em", "ode", "heun"]: row = f" {name:<8}" for n in n_list: row += f"{res[(name, n)]['dspread']:>13.4f}" print(row) print() print("=" * 74) print("为什么 DDIM 大步长扛得住、Euler 扛不住:看每步跨过的 L 有多大") print("=" * 74) print(f" {'N':>6} {'L_max':>10} {'L_mean':>10} {'e^(L/2) 与 1+L/2 的相对误差':>30}") for r in stride_stats(): print(f" {r['n']:>6} {r['L_max']:>10.4f} {r['L_mean']:>10.4f} " f"{r['euler_err'] * 100:>27.2f}%") print(f"\n 基线:真实样本 mean log p_0 = {base['logp']:.4f}," f"TV = {base['tv']:.4f},spread err = {base['spread']:.4f}") print(f" 指标自身噪声(另一份独立真实样本对基线的偏移):" f"logp {noise['logp']:+.4f} / TV {noise['tv']:+.4f} / " f"spread {noise['spread']:+.4f}") print() print(" 按 NFE(score 评估次数)对齐再看一遍——Heun 每步算两次:") print(f" {'采样器':<8} {'N':>6} {'NFE':>6} {'dlogp':>11} {'dspread':>11}") for name, n in [("ddpm", 1000), ("ddpm", 200), ("ddim", 50), ("ddim", 100), ("em", 50), ("ode", 50), ("heun", 25), ("heun", 50), ("ode", 200)]: r = res[(name, n)] print(f" {name:<8} {n:>6} {r['nfe']:>6} {r['dlogp']:>11.4f} " f"{r['dspread']:>11.4f}") print() print("=" * 74) print("关键对照:50 步能追上 1000 步吗(主指标看 dlogp,0 = 与真实样本一致)") print("=" * 74) for name, n in [("ddpm", 1000), ("ddim", 50), ("heun", 50), ("ode", 50), ("em", 50), ("ddpm", 50)]: r = res[(name, n)] print(f" {name:<6} N={n:<5} NFE={r['nfe']:<5} " f"dlogp={r['dlogp']:+.4f} dTV={r['dtv']:+.4f} " f"dspread={r['dspread']:+.4f} MMD^2={r['mmd2']:+.2e}") print() print("=" * 74) print("消融:只改 eta(注入多少噪声),其余完全不动,N=50") print("=" * 74) print(f" {'eta':>6} {'dlogp':>12} {'dTV':>12} {'MMD^2':>13}") for r in ablation_noise(): print(f" {r['eta']:>6.2f} {r['dlogp']:>12.4f} {r['dtv']:>12.4f} " f"{r['mmd2']:>13.2e}") print() print("=" * 74) print("换余弦调度再跑一遍(N=50)") print("=" * 74) res_c, _ = sweep(n_list=(50,), schedule="cosine") print(f" {'采样器':<8} {'dlogp(linear)':>15} {'dlogp(cosine)':>15}") for name in ["ddpm", "ddim", "em", "ode", "heun"]: print(f" {name:<8} {res[(name, 50)]['dlogp']:>15.4f} " f"{res_c[(name, 50)]['dlogp']:>15.4f}") print() print(f" (参照:MMD^2 噪声地板 = {floor:.3e},RBF 带宽 = {bw:.3f};" f"真实样本 mean log p_0 = {base['logp']:.4f})") if __name__ == "__main__": main()
2026年09月27日
4 阅读
0 评论
0 点赞
2026-09-26
AIGC 基本功|VAE 结构与训练目标-VAE
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. 最小可用理解 三句话: 结构:编码器把 $x$ 压成 $2d$ 个数($d$ 个均值 $\mu$、$d$ 个对数方差 $\log\sigma^{2}$),重参数化采出一个 $z$,解码器从 $z$ 重建 $x$。训练目标是 ELBO——本质上是「重建质量」减「每个潜变量维度花掉的 KL 预算」。 KL 项既是正则,也可以看作信息预算。一维潜变量要花掉多少 KL,就必须换回足够的重建收益,否则最优解就是关掉这一维。实测里 $\beta = 10^{-3}$ 时模型正好活 4 个维度(合成数据的真实因子数就是 4),$\beta$ 收到 0.3 只剩 3 个,$\beta = 1$ 一个不剩。维度是一个个死的,不是一起死。 潜空间的尺度是 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. 经典论文脉络 Kingma & Welling, 2013, Auto-Encoding Variational Bayes(arXiv:1312.6114) — VAE 本体。两个贡献撑起了后面所有工作:重参数化技巧(把采样挪出计算图)和对角高斯先验下 KL 的闭式解(3.2 节那个式子)。 Higgins 等, ICLR 2017, β-VAE: Learning Basic Visual Concepts with a Constrained Variational Framework — 把 KL 权重从固定的 1 变成可调旋钮,$\beta>1$ 换解耦表征。副产品是把「后验坍缩」推到了台前:$\beta$ 一大,重建就塌。 van den Oord 等, 2017, Neural Discrete Representation Learning(arXiv:1711.00937) — VQ-VAE。用码本查表替换连续高斯采样,潜空间变成离散索引,3.2 节那个 KL 项整个消失了。 Esser 等, CVPR 2021, Taming Transformers for High-Resolution Image Synthesis(arXiv:2012.09841) — VQGAN。给 VQ 自编码器加上感知损失和对抗损失,把重建质量推到可商用,第一次让「学到的潜空间 + 自回归」在 1024² 上真正可用。 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. 延伸阅读 前置:变分下界与重参数化 —— ELBO 的完整推导和重参数化为什么可导,本篇只用结论。 同方向:视频 VAE 的常见 loss 组合 —— 真实训练里 $\beta$ 不是单独出现的,它要和 L1、LPIPS、GAN 四项一起配比。 数值细节:混合精度与数值稳定性 —— logvar 截断那类技巧在 fp16 下会变得更关键。 下一步:潜空间扩散与 Stable Diffusion 架构会讲 $\text{Var}(z)$ 和信噪比这条线在 UNet 采样里怎么具体表现出来;「离散化表征:VQ-VAE 与 VQGAN」讲另一条绕开后验坍缩的路线。 附录:完整代码 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())
2026年09月26日
3 阅读
0 评论
0 点赞
2026-09-25
AIGC 基本功|变分下界与重参数化-ELBO
变分下界与重参数化 所属方向:数学基础 | 难度:入门 | 前置知识:无(会求导、会算期望即可) 关键词:变分推断、ELBO、重参数化、KL 散度、后验坍缩、β-VAE、IWAE、得分函数估计量 01. 为什么需要它 先摆三组在同一台机器上跑出来的数字,都出自文末附录,可以自己复现。 第一组:KL 项掉到 0,潜变量整条死掉。 固定解码器噪声 σ²=0.25,让四个独立数据方向分别具有方差 λ=4.0、1.0、0.25、0.0625,只改 KL 项的权重 β,看看最优解长什么样: 数据方向方差 λ β=0.5 β=1.0 β=2.0 β=4.0 死亡阈值 β* 4.0 MSE 0.125 存活 MSE 0.250 存活 MSE 0.500 存活 MSE 1.000 存活 16.00 1.0 MSE 0.125 存活 MSE 0.250 存活 MSE 0.500 存活 MSE 1.000 死亡 4.00 0.25 MSE 0.125 存活 MSE 0.250 死亡 死亡 死亡 1.00 0.0625 死亡 死亡 死亡 死亡 0.25 死亡那一格的具体解是 a=0、s=1、KL=0:编码器把均值直接输出 0、方差输出 1,等于彻底放弃这个维度。此时重建误差精确等于 λ——也就是把 x 全猜成 0 的水平,这个维度一点信息都没传。问题在于看 loss 曲线你只会看到「KL 顺利降到 0,总损失还在降」,很难意识到这是失败而不是收敛。 第二组:把 ELBO 当成似然来监控,会被方差骗。 在一个 D=6、K=3 的高斯线性模型上,真实 log p(x) = -5.434331,解析 ELBO = -7.419767,两者差 1.985436。但用训练时真正用的那个单样本蒙特卡洛估计量去估 ELBO,重复 4000 次,标准差是 4.308——比它和 log p(x) 之间的差距还大一倍多,结果有 40.7% 的采样值直接越过了 -5.434 这条「上界」。如果你在 tensorboard 上画这条线当似然看,会得出完全错误的结论。 第三组:不用重参数化,梯度方差能大到没法训练。 同一个模型上估 ∇ E_q[log p(x|z)],重参数化的梯度方差和是 377.55,得分函数估计量是 3043.11,差 8.1 倍。听起来还能忍?把数据维度 D 从 8 拉到 2048: D 重参数化方差 得分函数方差 比值 8 2.08 2 469 1 188× 64 1.98 104 328 52 630× 512 1.84 7 579 291 4 121 907× 2048 1.84 114 230 222 61 916 893× 重参数化的方差几乎不随 D 动,得分函数估计量按 D² 往上冲,到 D=2048(一张 45×45 的灰度图而已)比值是 6191 万倍。这不是「慢一点」的差别,是「能不能训」的差别。 这三组数字背后是同一件事:我们想最大化的是 log p(x),但它算不动,只能换成它的一个下界 ELBO 来优化;而换成下界之后,「换掉了什么」「这个下界的估计量有多吵」「梯度怎么穿过采样」这三笔账必须自己算清楚。不搞清楚 KL 项的价,你就不知道 β 该往哪调;不搞清楚下界的松紧,你就不知道扩散模型的训练目标从哪来。 02. 最小可用理解 三句话讲完核心: 算不动的边际似然换成一个能算的下界:$p_\theta(x)$ 要对所有 z 积分,算不动,于是引入一个我们自己挑的分布 $q_\phi(z \mid x)$,用 Jensen 不等式把优化目标换成 ELBO,它只需要在 q 下求期望,采样就能算。 下界和真值之间差的那一坨,正好是 q 和真实后验的 KL:$\log p_\theta(x) = \mathcal{L}(\theta, \phi) + \mathrm{KL}(q_\phi(z \mid x) \,\|\, p_\theta(z \mid x))$。所以「下界有多松」等价于「q 离真后验有多远」,而不是「q 离先验有多远」——后者的 KL 是 ELBO 里的一个加数,是另一回事。 要让这个下界能被反向传播,就把随机性从梯度路径上挪走:写成 $z = m_\phi + s_\phi \odot \epsilon$,$\epsilon$ 与参数无关,采样这一步变成确定性变换,梯度顺着 z 一路传回 $m_\phi$ 和 $s_\phi$。 03. 数学推导 3.1 边际似然为什么算不动 生成模型的设定很简单:先从一个固定先验里采隐变量 $z \sim p(z)$(通常是标准正态),再用解码器生成观测 $x \sim p_\theta(x \mid z)$。要拟合数据,目标是对数边际似然: $$\log p_\theta(x) = \log \int p_\theta(x \mid z) \, p(z) \, dz$$ 麻烦全在这个积分上。$p_\theta(x \mid z)$ 是个神经网络,塞进积分里没有任何闭式解;用数值积分的话,z 是 K 维,网格点数随 K 指数爆炸,K=32 就已经不可能。而 K 小了模型表达能力又不够——这正是我们不愿意接受的取舍。 3.2 换个能算的目标 既然积分算不动,就绕开它。引入在目标分布支撑上为正、满足相关可积条件的以 x 为条件的分布 $q_\phi(z \mid x)$(后面简称 $q$),把被积函数乘一个 $q/q$: $$\log p_\theta(x) = \log \int q_\phi(z \mid x) \, \frac{p_\theta(x, z)}{q_\phi(z \mid x)} \, dz = \log \mathbb{E}_{q_\phi}\!\left[\frac{p_\theta(x, z)}{q_\phi(z \mid x)}\right]$$ log 是凹函数,Jensen 不等式给出 $\log \mathbb{E}[Y] \ge \mathbb{E}[\log Y]$,于是 $$\log p_\theta(x) \ge \mathbb{E}_{q_\phi}\!\left[\log \frac{p_\theta(x, z)}{q_\phi(z \mid x)}\right] \equiv \mathcal{L}(\theta, \phi)$$ 右边就是 ELBO(Evidence Lower Bound)。它好算:$\log p_\theta(x, z) = \log p_\theta(x \mid z) + \log p(z)$ 两项都能直接求值,期望用采样估计就行。代价是我们不再直接优化 log p(x),而是优化它的一个下界。 3.3 差的那一项到底是什么 这是全文最关键的一步,也是最容易被跳过的一步。把 $\log p_\theta(x)$ 写成对 q 的期望(它和 z 无关,所以这么写是恒等的),再硬塞一个 $\log \frac{q}{q}$ 进去: $$\log p_\theta(x) = \mathbb{E}_{q}\!\left[\log p_\theta(x)\right] = \mathbb{E}_{q}\!\left[\log \frac{p_\theta(x, z)}{p_\theta(z \mid x)}\right] = \mathbb{E}_{q}\!\left[\log \frac{p_\theta(x, z)}{q_\phi(z \mid x)}\right] + \mathbb{E}_{q}\!\left[\log \frac{q_\phi(z \mid x)}{p_\theta(z \mid x)}\right]$$ 第一项正是 ELBO,第二项按定义就是 KL 散度。合起来: $$\log p_\theta(x) = \mathcal{L}(\theta, \phi) + \mathrm{KL}\big(q_\phi(z \mid x) \,\|\, p_\theta(z \mid x)\big)$$ KL 恒非负,所以 $\mathcal{L} \le \log p_\theta(x)$,和 Jensen 的结论一致——但这一版多给了一个信息:等号成立当且仅当 q 等于真实后验。下界的松紧完全由 q 的质量决定,跟别的都没关系。 3.4 把 ELBO 拆成两项看 把 $\log p_\theta(x, z)$ 展开,ELBO 可以写成更有物理含义的形式: $$\mathcal{L}(\theta, \phi) = \mathbb{E}_{q_\phi}\!\left[\log p_\theta(x \mid z)\right] - \mathrm{KL}\big(q_\phi(z \mid x) \,\|\, p(z)\big)$$ 第一项是重建项:从 q 里采 z,解码器能不能还原出 x。第二项是先验对齐项:q 别离先验太远——因为生成时我们是从先验采 z 的,如果 q 把 z 放到了先验覆盖不到的地方,生成阶段就对不上了。 这里有个必须记住的区分,后面第 08 节还会回到它: $\mathrm{KL}(q \,\|\, p(z))$ —— 出现在 ELBO 里,是要被最小化的一项,物理含义是「编码分布别跑太偏」。 $\mathrm{KL}(q \,\|\, p(z \mid x))$ —— 不出现在 ELBO 里,是下界的缝隙,物理含义是「q 离真后验还差多少」。 在我们那个 D=6 的例子里,这两个数是 3.364364 和 1.985436,不是一回事,也不是简单的包含关系。β-VAE 显式改的是先验 KL 的权重,但重新训练后 q 与生成模型都会变化,所以真实后验差距也会变化,且不保证变小。 顺带一个很容易漏掉的推论:把 KL 乘上 β 之后得到的目标 $$\mathcal{L}_\beta = \mathbb{E}_{q_\phi}\!\left[\log p_\theta(x \mid z)\right] - \beta \cdot \mathrm{KL}\big(q_\phi(z \mid x) \,\|\, p(z)\big)$$ 当 $\beta\ge1$,仍有 $\mathcal L_\beta=\mathcal L-(\beta-1)\mathrm{KL}(q\Vert p)\le\mathcal L\le\log p_\theta(x)$,所以它仍是下界,只是更松;$0<\beta<1$ 时则不再保证下界。不同 β 的目标包含不同惩罚,不能直接当成同一种似然估计比较。 3.5 梯度怎么穿过采样 ELBO 对解码器参数 $\theta$ 的梯度没问题,重建项直接可导。麻烦在对 $\phi$ 的梯度:期望的分布本身依赖于 $\phi$,而采样操作不可导。设 $f(z) = \log p_\theta(x \mid z)$,要估的是 $\nabla_\phi \mathbb{E}_{q_\phi}[f(z)]$。 方法一:得分函数估计量(REINFORCE)。 直接把导数挪进期望: $$\nabla_\phi \mathbb{E}_{q_\phi}[f(z)] = \mathbb{E}_{q_\phi}\!\left[f(z) \, \nabla_\phi \log q_\phi(z \mid x)\right]$$ 在可交换微分与积分、分布支持集不随参数变化等常见正则条件下,这个式子成立;它不要求 f 对 z 可导,也不要求 z 连续——代价是方差极大,因为整个 f 的量级都被乘进了梯度里。 方法二:重参数化。 把 z 写成参数的确定性函数外加一个与参数无关的噪声: $$z = m_\phi(x) + s_\phi(x) \odot \epsilon, \quad \epsilon \sim \mathcal{N}(0, I)$$ 于是期望可以改写成对 $\epsilon$ 的期望,导数直接进去: $$\nabla_\phi \mathbb{E}_{q_\phi}[f(z)] = \mathbb{E}_{\epsilon}\!\left[\nabla_\phi f\big(m_\phi(x) + s_\phi(x) \odot \epsilon\big)\right] = \mathbb{E}_{\epsilon}\!\left[\nabla_z f(z) \cdot \nabla_\phi z\right]$$ 关键在于 $\nabla_z f$ 用到了 f 对 z 的局部形状,相当于「知道往哪个方向挪 z 会让 f 变大」,这是方法一完全没有利用的信息。 方差为什么差这么多——一个具体的机制。 在高斯解码器下 $$f(z) = -\frac{D}{2}\log(2\pi\sigma^2) - \frac{\|x - Wz - b\|^2}{2\sigma^2}$$ 第一项与 z 无关,是个常数。它对 $\nabla_z f$ 的贡献恒为 0,重参数化天然看不见它;但方法一要把整个 f(包括这个常数)乘上 $\nabla_\phi \log q$,常数按平方进方差。$D=2048$、$\sigma^2=1$ 时这个常数是 $-\frac{2048}{2}\log 2\pi \approx -1877$,平方之后就是 350 万量级——这正是第三组数字里那个 6191 万倍的来源。 实验也验证了这一点:给方法一减掉一个「预言机 baseline」(把 f 换成 $f - \mathbb{E}_q[f]$,常数的均值被减掉,期望不变),方差比从 61 916 893× 掉到 5.3×。所以差距主要来自那个常数项,不是来自「采样本身」。 3.6 下界能有多紧:IWAE 单样本 ELBO 的缝是 $\mathrm{KL}(q \,\|\, p(z \mid x))$。要缩缝有两条路:把 q 变强(换更灵活的后验族),或者换一个更紧的界。后者的经典做法是 IWAE 的 k 样本界: $$L_k = \mathbb{E}_{z_1 \dots z_k \sim q_\phi}\!\left[\log \frac{1}{k} \sum_{l=1}^{k} \frac{p_\theta(x, z_l)}{q_\phi(z_l \mid x)}\right]$$ 注意顺序:先对 k 个重要性权重取平均,再取 log。单样本 ELBO 是「log 再取平均」(每个样本的 log 比值取期望),IWAE 是「平均再 log」,Jensen 保证后者更紧。可以证明 $\log p_\theta(x) \ge L_{k+1} \ge L_k \ge L_1 = \mathcal{L}$,在重要性权重满足支持覆盖与可积性等条件下,$k \to \infty$ 时收敛到 $\log p_\theta(x)$。 「换个顺序就更紧」这件事直觉上可以这么理解:单样本 ELBO 每次只看一个 z,好坏全押在它身上;IWAE 一次看 k 个,其中只要有一个 z 的 $p_\theta(x, z)/q_\phi(z \mid x)$ 特别大,平均权重就被拉上去,而 log 是凹函数,对这个「运气好」的样本惩罚得比线性小。所以 k 越大,越有机会碰到好样本,界越紧——本质上是用更多的采样换取更少的方差,跟蒙特卡洛里加样本降噪是同一回事,但这里降的是偏差而不是方差。 代价是:梯度不再是对单个样本的简单求和,k 个权重互相耦合(每个权重的梯度都带上了其他权重的归一化因子),总计算量随 k 线性上涨;编码器梯度的信噪比还可能随 k 增大而下降(见 Rainforth 等,2018),这反而不利于训练编码器——工程上这是个明确的取舍,不是免费的午餐。 这张图要看什么:左图四根柱子——ELBO(-7.420)加 KL(q‖p(z|x))(1.985)正好等于 log p(x)(-5.434),而 KL(q‖p(z))(3.364)是另一根完全不同的柱子,别混。中图是单样本蒙特卡洛 ELBO 的 4000 次采样分布,标准差 4.31,虚线是真实上界 -5.434,40.7% 的样本落在它右边。右图是 IWAE 的 k 扫描,k 从 1 涨到 500,缝隙从 2.049 缩到 0.0026——同样是用重要性采样,换个顺序就差三个数量级。 04. 代码实现 4.1 恒等式核对:先造一个能算到机器精度的模型 要验证「差的那一项到底是什么」,必须三样东西都有解析解:边际似然、真实后验、ELBO。高斯线性模型满足这一点:$z \sim \mathcal{N}(0, I_K)$,$x \mid z \sim \mathcal{N}(Wz + b, \sigma^2 I_D)$。边际仍是高斯,后验是共轭高斯,ELBO 也能闭式算。 import numpy as np D, K, SIGMA2 = 6, 3, 0.35 ** 2 W = np.random.default_rng(0).normal(0, 1, size=(D, K)) * 0.9 b = np.random.default_rng(1).normal(0, 0.3, size=(D,)) def true_posterior(x): """p(z|x) 的闭式解:precision = I + W^T W / sigma^2。""" prec = np.eye(K) + W.T @ W / SIGMA2 cov = np.linalg.inv(prec) return cov @ (W.T @ (x - b) / SIGMA2), cov def elbo_analytic(x, m, s): """ELBO 的解析值。注意 ||x - Wz - b||^2 对 q 求期望会多出一项 trace。""" var = s ** 2 resid = x - (W @ m + b) recon = -0.5 * (D * np.log(2 * np.pi * SIGMA2) + (resid @ resid + np.trace(W.T @ W @ np.diag(var))) / SIGMA2) kl = 0.5 * np.sum(var + m ** 2 - 1.0 - np.log(var)) return float(recon - kl), float(recon), float(kl) def kl_gauss_gauss(m, s, mean2, cov2): """KL(N(m, diag(s^2)) || N(mean2, cov2)),两个高斯的闭式。""" var, diff = s ** 2, mean2 - m cov2_inv = np.linalg.inv(cov2) _, logdet2 = np.linalg.slogdet(cov2) return 0.5 * (np.trace(cov2_inv @ np.diag(var)) + diff @ cov2_inv @ diff - K + logdet2 - np.sum(np.log(var))) 重建项里那个 np.trace(W.T @ W @ np.diag(var)) 是最容易漏掉的一项:对 $z$ 求期望时,$\|x - Wz - b\|^2$ 里的 $Wz$ 项会因为 z 的随机性多出一份方差贡献。漏了它,ELBO 就不是下界了。 关键一步是故意把 q 指定错:取真实后验的均值再加扰动,协方差只留对角线并缩小 20%,让 q 比真后验更自信且忽略相关性。这样 KL(q‖p(z|x)) 严格大于 0,缝才看得见。跑起来: [1] 解析三项 E_q[log p(x|z)] = -4.055403 KL(q || p(z)) = 3.364364 ELBO = -7.419767 KL(q || p(z|x)) = 1.985436 ELBO + KL(q||p(z|x)) = -5.434331 log p(x) = -5.434331 |误差| = 8.882e-16 恒等式对到 1e-15。同时注意:KL(q‖p(z))=3.364 比 KL(q‖p(z|x))=1.985 还大——这两个量的大小关系没有必然规律,别用其中一个去猜另一个。 4.2 两种梯度估计量,方差实测 同一个模型上,$\nabla_m$ 和 $\nabla_\ell$($\ell = \log s$)的解析梯度都能写出来,所以「谁对谁错」有标准答案,剩下的差别纯粹是方差: def sample_reparam(x, m, s, W, b, sigma2, rng, n): """重参数化:z = m + s ⊙ ε,梯度经 z 反传。""" eps = rng.normal(size=(n, m.size)) z = m + s * eps score_z = (x - (z @ W.T + b)) @ W / sigma2 # ∇_z log p(x|z) return score_z, score_z * s * eps # ∂z/∂m = 1, ∂z/∂ℓ = s ⊙ ε def sample_score(x, m, s, W, b, D, sigma2, rng, n, baseline=None): """得分函数:∇ E[f] = E[f · ∇ log q],f = log p(x|z)。""" eps = rng.normal(size=(n, m.size)) z = m + s * eps r = x - (z @ W.T + b) f = -0.5 * (D * np.log(2 * np.pi * sigma2) + np.einsum("nd,nd->n", r, r) / sigma2) if baseline is not None: f = f - baseline return f[:, None] * (eps / s), f[:, None] * (eps ** 2 - 1.0) 两行代码的差别就是全部:重参数化用的是 $\nabla_z f$,得分函数用的是 $f$ 本身。跑 40000 次独立采样: 真实梯度(解析) ∇_m = [ 0.4228 4.9273 -0.6448] ∇_ℓ = [-1.9319 -2.4426 -3.8296] 估计量对比(40000 次独立采样,方差 = 6 个参数分量方差之和) 重参数化 方差和= 377.55 均值相对误差= 0.41% 达到5%需 n≈3068 得分函数(无 baseline) 方差和= 3043.11 均值相对误差= 3.95% 达到5%需 n≈24722 得分函数(预言机 baseline) 方差和= 1586.45 均值相对误差= 3.09% 达到5%需 n≈12889 方差比:得分函数 / 重参数化 = 8.1× 加 baseline 之后 = 4.2× 三个估计量的均值都对(无偏),差别全在方差。换算成「达到 5% 相对误差需要多少样本」,是 3068 对 24722。 这张图要看什么:左图是 $\nabla_m$ 第 0 个分量(真值 0.4228)在三种估计量下的分布——均值都压在真值附近,标准差分别是 5.7、16.9、12.0,尾巴长度差一个量级。右图的双对数坐标里,重参数化那条线几乎是平的(方差 2.08 → 1.84),得分函数那条严格贴着 D² 参考线往上走;而把蓝色点(加了 baseline)和橙色点对比,比值从 1188× 一路到 6191 万×,加完 baseline 却稳定在 5~7×——说明炸掉的部分是那个常数项,不是采样噪声。 4.3 β 加权:把维度死掉的过程解出来 这一节的目标是要一个能解到全局最优的玩具模型,这样「维度死了」是解析结论而不是训练运气。构造:每个坐标独立,$x_j \sim \mathcal{N}(0, \lambda_j)$,$q(z_j \mid x_j) = \mathcal{N}(a_j x_j, s_j^2)$,$p(x_j \mid z_j) = \mathcal{N}(b_j z_j, \sigma^2)$。每个维度就是一份独立副本,逐个维度单独求最优即可: SIGMA2 = 0.25 # 解码器观测噪声方差 def kl_of(a, lam, s): return 0.5 * (a ** 2 * lam + s ** 2 - 1.0 - np.log(s ** 2)) def recon_mse(a, b, lam, s): """E_{x,z}[(x - b z)^2],重建误差的期望。""" return lam * (1.0 - b * a) ** 2 + b ** 2 * s ** 2 def objective(p, lam, beta): a, b, ell = p s = np.exp(ell) return -recon_mse(a, b, lam, s) / (2 * SIGMA2) - beta * kl_of(a, lam, s) 目标对 $(a, b, \ell)$ 的梯度全部解析,脚本里再用中心差分核对一遍(实测最大偏差 5.5e-10),避免推错。优化必须多起点——$(a, b) = (0, 0)$ 就是「维度死亡」解,单起点很容易掉进去出不来。 跑完 β 扫描,把最优解和闭式预测放在一起对: λ β s²实测 βσ²/λ a²λ实测 1-s² MSE实测 βσ² 1.0000 1.00 0.25000 0.25000 0.75000 0.75000 0.25000 0.25000 1.0000 2.00 0.50000 0.50000 0.50000 0.50000 0.50000 0.50000 1.0000 3.00 0.75000 0.75000 0.25000 0.25000 0.75000 0.75000 0.2500 0.50 0.50000 0.50000 0.50000 0.50000 0.12500 0.12500 最大偏差 = 5.53e-03 → 存活时最优解确实落在闭式上 λ=4.0000 理论死亡阈值 β* = λ/σ² = 16.00 λ=1.0000 理论死亡阈值 β* = λ/σ² = 4.00 λ=0.2500 理论死亡阈值 β* = λ/σ² = 1.00 λ=0.0625 理论死亡阈值 β* = λ/σ² = 0.25 存活时最优解落在闭式 $s^2 = \beta\sigma^2/\lambda$、$a^2\lambda = 1 - s^2$、$\mathrm{MSE} = \beta\sigma^2$ 上。这个闭式有个反直觉的推论:存活时的重建误差只由 β 和解码器噪声决定,跟这个维度携带多少信息 λ 无关——λ 只决定「这个维度值不值得用」。 05. 工业级实现对照 最小实现里 KL 是两个高斯的闭式,工业代码也是闭式,但有几个地方长得不一样,值得逐个说明为什么。 1. 官方最小实现:pytorch/examples。 pytorch/examples 的 vae/main.py 里,loss_function 就是标准的 BCE 重建加闭式 KL,没有任何 β(以 2026-09 时的实现为准): def loss_function(recon_x, x, mu, logvar): BCE = F.binary_cross_entropy(recon_x, x.view(-1, 784), reduction='sum') # 0.5 * sum(1 + log(sigma^2) - mu^2 - sigma^2) KLD = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp()) return BCE + KLD 和本文 3.4 节那一项完全对得上:$\mathrm{KL} = \frac{1}{2}\sum(\mu^2 + s^2 - 1 - \log s^2)$,代码里写成 -0.5 * sum(1 + logvar - mu^2 - exp(logvar)),logvar 就是 $\log s^2$。要加 β 得自己动手,官方示例没有。 2. 为什么参数化 logvar 而不是 s。 方差必须为正,直接学 s 要加约束;学 $\log s^2$ 则值域是全体实数,网络怎么输出都合法。diffusers 的 DiagonalGaussianDistribution 还额外做了截断(以 2026-09 时的实现为准): self.mean, self.logvar = torch.chunk(parameters, 2, dim=1) self.logvar = torch.clamp(self.logvar, -30.0, 20.0) self.std = torch.exp(0.5 * self.logvar) self.var = torch.exp(self.logvar) clamp 到 [-30, 20] 是为了数值稳定:$\exp(20)$ 约为 4.85 亿,fp32 可表示但 fp16 已无法表示;下界 -30 避免过小方差与过大负 logvar。还需核对计算 dtype,截断本身不是所有精度下的安全保证。 3. 采样就是重参数化,一行代码。 同一个类里: def sample(self, generator=None): sample = randn_tensor(self.mean.shape, generator=generator, device=self.parameters.device, dtype=self.parameters.dtype) x = self.mean + self.std * sample return x 就是本文 3.5 的 $z = m + s \odot \epsilon$。mode() 直接返回均值——一些确定性编码场景使用 mode(),但 diffusers img2img 默认使用后验 sample(generator=...),训练和推理均可采样。 4. KL 的求和维度与归一化。 diffusers 里 kl() 写成(以 2026-09 时的实现为准): return 0.5 * torch.sum(torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar, dim=[1, 2, 3]) 对通道、高、宽三个维度求和,得到的是每个样本一个标量,然后训练脚本再对 batch 取平均。这个区别很实际:如果对全部分量求和再除以元素总数,等价于把每个维度的 KL 权重除以 $C \times H \times W$——潜变量维度一多,KL 项的有效权重会相应减小,可能使先验对齐变弱、尺度漂移;过强 KL 才更直接推动后验坍缩。 5. 训练时 KL 根本不在模型里。 值得注意的是 AutoencoderKL 的模型文件里没有 kl_loss 这个东西,它只通过 encode() 返回 DiagonalGaussianDistribution 后验,KL 由外部训练脚本算。这是工程上的职责划分:模型只负责给出分布,损失函数怎么组合(KL 权重、perceptual loss、对抗损失)交给训练配置。 6. 潜空间的缩放。 SD1.x 常见 scaling_factor=0.18215 是乘到原始 latent 上的系数,对应原始标准差约 $1/0.18215\approx5.49$,缩放后才接近 1,不能把 0.18 当成原始潜空间尺度。它来自特定 VAE 与扩散训练的尺度约定,不由 ELBO 唯一决定。 7. 最小实现和工业实现的差异。 工业代码增加 logvar 截断、dtype 控制、分布采样接口及 latent 尺度约定。KL 逐维求和来自概率定义,batch 平均来自训练目标;它们不是任意工程细节。sample() 与 mode() 都有推理用途,要依具体管线选择。 06. 代价与边界 省了什么。 把一个对 K 维积分的不可解问题,变成了「采样 + 两个可导项」,梯度能用标准反向传播算。这是 VAE、扩散模型、以及一大票潜变量模型能训练起来的全部前提。 赔了什么,四条。 一是下界不是似然。 优化 ELBO 不等于优化 log p(x),中间隔着 KL(q‖p(z|x))。在固定生成模型下,这个缝受 q 族表达能力与实际优化结果共同影响:均值场对角高斯拟合不了有相关性的真实后验,缝就永远在。本例里缝是 1.985 nats,在一个 log p(x) 只有 -5.43 的小模型上,数值上差约 36%,但连续对数密度受坐标单位影响,不宜把这个比例当作通用误差尺度。 二是 KL 的方向有偏好。 $\mathrm{KL}(q\Vert p(z\mid x))$ 是 reverse KL:当 $q$ 把概率放到真实后验很低的地方,代价很大,因此在受限的单峰近似族下常表现为 mode-seeking,可能只覆盖一个峰;反向的 $\mathrm{KL}(p\Vert q)$ 则倾向覆盖目标的质量,即 mass-covering。VAE 的模糊不能单归因于这个方向,逐像素重建目标、解码器分布与压缩瓶颈都有关。 三是推断被参数化摊平(amortized)之后又多了一层误差。 经典变分推断对每个 x 单独优化一组变分参数;VAE 用一个共享的编码器网络去输出所有 x 的 $m_\phi(x)$ 和 $s_\phi(x)$。这一步是为了快——测试时一次前向就得到后验,不用重新迭代——但它意味着 q 的可行域被限制在「神经网络能表达的那些分布」里。即便每个 x 单独看,最优的对角高斯后验就在那儿,共享编码器也可能一辈子到不了。所以总的缝其实是两笔账叠起来的:函数族的近似误差加上摊销误差。这也解释了为什么给编码器加容量在有些任务上能明显提 ELBO——在固定生成模型时,这改善的是近似推断。 四是加权 KL 会直接杀死维度。 这是第 4.3 节那张表的完整结论: λ=4.0000 在 β≤12.0 内都存活 λ=1.0000 从 β=4.0 起死亡(β=1 时 KL=0.6931) λ=0.2500 从 β=1.0 起死亡 λ=0.0625 从 β=0.25 起死亡 死亡阈值是 $\beta^\ast = \lambda / \sigma^2$。读出这个式子的含义:一个维度要活下来,它携带的信息量 λ 必须盖过「KL 的价」β 乘上「解码器噪声」σ²。β 翻倍,能活下来的维度门槛就翻倍;解码器越准(σ² 越小),越多的维度能活。这仅解释线性高斯模型的阈值。强大的自回归解码器可通过其他条件预测数据而忽略 z,不能简单等同于把固定高斯观测方差 σ² 调小。 这张图要看什么:左图四条曲线是不同 λ 下重建 MSE 随 β 的变化,实心点表示维度存活、空心方块表示已死——注意存活段的 MSE 就是 $\beta\sigma^2$ 这条直线,跟 λ 无关;而死掉的段 MSE 平在 λ 上不再变化。右图是 $(\lambda, \beta)$ 平面上的相图,斜线是 $\beta = \lambda/\sigma^2$,线右上方全死、左下方全活。调 β 之前先看一眼这个平面:你真正要判断的是「这条线上方还有多少维度」。 缓解手段:free bits。 把每组 KL 换成 $\max(\mathrm{KL},C)$,是在 KL 小于 C 时去掉进一步压低它的梯度,给重建项使用这部分容量的机会;这不是保证至少传 C nats 的硬约束。下面的线性实验确实在边界附近找到更好的重建,但一般神经网络不保证被救活。实测(β=4): λ C=0(纯 β) C=0.05 C=0.2 1.0000 a=0.000 KL=0.000 MSE=1.000 a=-0.308 KL=0.050 MSE=0.905 a=-0.574 KL=0.200 MSE=0.670 0.2500 a=0.000 KL=0.000 MSE=0.250 a=-0.617 KL=0.050 MSE=0.226 a=-1.148 KL=0.200 MSE=0.168 λ=1.0 的维度被救回来了(MSE 1.000 → 0.670),λ=0.25 的也从 0.250 降到 0.168。代价是放松了对先验 KL 的惩罚;是否改善或损害解耦需要独立评估——这是一笔明码标价的交易。 什么时候不该用。 归一化流和自回归模型在相应建模假设下可计算精确似然,不一定需要变分下界;VQ 模型可精确计算离散 token 序列的自回归概率,但一般仍不能精确边缘化得到像素似然。ELBO 也不是训练高维生成模型的唯一途径,score matching、流匹配和对抗训练是其他路线。同一数据、同一似然约定下标准 ELBO 可比较为下界,但 q 的质量影响松紧,不能直接据此断言真实似然或感知质量的排序。 07. 经典论文脉络 Kingma & Welling, 1312.6114(2013)——VAE 原文。贡献是把「变分推断 + 重参数化 + 神经网络编码器」拼成一个能用 SGD 训的东西,并给出 SGVB 估计量。它同时确立了沿用至今的 loss 形式:重建项减 KL。本文 3.2~3.5 节基本是这篇的复述。 Rezende, Mohamed & Wierstra, 1401.4082(2014)——几乎同期独立提出的重参数化,论文里叫 stochastic backpropagation。贡献是把这个方法从「VAE 的一个技巧」推广成「任何可微分概率模型上的通用推断方法」,并系统讨论了高斯之外的分布族怎么处理。想理解重参数化的适用边界(哪些分布能做、哪些只能退回得分函数),这篇比 VAE 原文讲得更清楚。 Burda, Grosse & Salakhutdinov, 1509.00519(2015)——IWAE。贡献就是本文 3.6 那个 $L_k$:把「log 再平均」换成「平均再 log」,得到一个随 k 单调变紧的界,并证明了收敛性。它澄清了一个当时普遍的误解——「多采几个样本只是降方差」,实际上换的是界本身。 Higgins et al., ICLR 2017——β-VAE。给 KL 项加权重 β > 1,换来更好的解耦表征。本文第 4.3 节那张表就是它的代价面:β 每翻一倍,重建误差也翻一倍,且维度按 $\beta^\ast = \lambda/\sigma^2$ 逐个死掉。 Bowman et al., 1511.06349(2015)——后验坍缩最早被认真对待的现场。用 VAE 做句子生成,强大的自回归解码器会直接忽略潜变量,KL 掉到 0。这篇提出的 KL annealing(训练初期把 β 从 0 慢慢涨到 1)至今还是最实用的缓解手段之一,效果依赖具体模型和优化过程,并非保证解决坍缩。 补充两篇:Kingma et al. 的 IAF(1606.04934)走的是另一条路——不改界,改 q,用可逆变换把后验族变强来直接缩小那个缝;Kingma & Welling 的综述(1906.02691)适合把上述脉络串起来通读。 08. 常见误解 误解一:「ELBO 里的 KL 就是 q 和真实后验的差距」。 ELBO 里减去的是 $\mathrm{KL}(q\Vert p(z))$;下界的差距是 $\mathrm{KL}(q\Vert p(z\mid x))$。调整 β 会通过训练改变 q 和生成模型,后者也可能变化,只是没有保证随前者一起下降。 误解二:「ELBO 是下界,所以蒙特卡洛估计值不会超过 log p(x)」。 下界性质是对期望成立的,不是对每个样本成立。本例单样本估计的标准差是 4.308,而缝只有 1.985,结果 40.7% 的样本越过了真实上界。看到自己的「ELBO」比之前算的 log p(x) 还大时,先别怀疑代码,这是正常的采样噪声。 误解三:「重参数化是为了让采样可微」。 更准确的说法是为了降低梯度估计量的方差。采样「不可微」这个表述本身就有问题:得分函数估计量里 $\nabla_\phi \log q_\phi$ 是对参数求导,完全可微,它对 f 连可导性都不要求。真正的区别是重参数化用上了 $\nabla_z f$ 这个局部信息,方差低几个数量级(D=2048 时差 6191 万倍)。 误解四:「KL 掉到 0 说明 KL 项优化到位了」。 这是后验坍缩的典型症状。本例 λ=1.0、β=4 时最优解就是 a=0、s=1、KL=0、MSE=1.000(正好等于 λ),编码器彻底放弃了这个维度。判断方法不是看 KL,而是看 $a^2\lambda$(Higgins 的 active unit 判据)——本文脚本里 active_unit() 就是干这个的。 误解五:「β 越大解耦越好,重建变差只是小代价」。 在线性模型中,最优 MSE 为 $\min(\beta\sigma^2,\lambda)$,连续增加后饱和;每维 KL 也连续趋零。离散的存活计数会出现台阶,但不意味着误差发生跳变。更强的 KL 约束还可能让所有维度关闭,不能只用 β 大小判断解耦质量。 09. 动手验证 三个脚本都只依赖 numpy,复制下来直接跑(完整代码见文末附录)。 python elbo_identity.py # 约 5 秒 python reparam_gradients.py # 约 15 秒 python beta_kl_weight.py # 约 3 分钟(多起点优化) 实验一:把 q 的均值和边际方差改成真后验的对应量,看残留差距。 改 elbo_identity.py 里的 make_q,把 var = np.diag(cov).copy() * 0.8 + 0.05 改成直接返回真实后验的对角(var = np.diag(cov).copy()),均值扰动也去掉。预期:KL(q‖p(z|x)) 会从 1.985 掉下来但不会掉到 0——因为对角 q 还是拟合不了真实后验的相关性。这个残留量就是「均值场假设的代价」,值得亲眼看一次。 实验二:给得分函数估计量换一个真实的 baseline。 reparam_gradients.py 里用的是预言机 baseline($\mathbb{E}_q[f]$ 的解析值),实战拿不到。把它换成一个滑动平均的 f(比如前 100 个样本的均值)再跑,预期方差比从 8.1× 降到接近 4.2× 的水平——这解释了为什么 REINFORCE 类算法里 baseline 是标配而不是优化项。 实验三:给 β 扫描加一个 free bits,看能救回几个维度。 beta_kl_weight.py 已经内置 free_bits 参数,把 main() 里的 optimize(lam, 4.0, free_bits=C) 的 C 从 0.05 调到 0.5 再跑。预期:这个线性实验中 C=0.5 时低方差维度也能使用 latent,KL 可能停在 C 附近;这不代表一般模型保证传满 C nats,也不能从 KL 数值单独判断解耦程度。这就是那笔交易的完整价格表。 10. 延伸阅读 按知识树的依赖关系,从这篇出发有三个方向: 往下(直接依赖本文):扩散模型的 SDE 视角——把 DDPM 的离散公式和连续 SDE 对上,会发现它的训练目标就是一层层套起来的 ELBO,本文 3.3 那个恒等式在那边换个马甲反复出现;视频 VAE 的潜空间与 loss 组合;其中 scaling_factor 的来历见 VAE 结构与训练目标。 往旁(同一个估计量,另一副马甲):策略梯度与 PPO 基础——本文 3.5 的得分函数估计量在 RL 里叫 REINFORCE,方差问题一模一样,baseline 的作用也一模一样。读过这篇再去看 PPO,会省掉一次重新理解。 已发布的相关篇目:视频 VAE 的常见 loss 组合——KL 之外还有哪些 loss 会加进来,以及它们各自的量级;FlashAttention 为什么不需要存下注意力矩阵 与 性能建模与 Profiling——本文没碰的工程侧。 完整的知识树见博客目录页 AIGC 基本功知识树,按依赖顺序排好了先修课。 附录:完整代码 09 节用到的脚本全文如下(elbo_identity.py、reparam_gradients.py、beta_kl_weight.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 elbo_identity.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """ELBO 恒等式核对:log p(x) = ELBO + KL(q || p(z|x))。 构造一个「后验有解析解」的模型,这样三样东西都能算到机器精度: * 边际似然 log p(x) —— 高斯线性模型,边际仍是高斯 * 真实后验 p(z|x) —— 共轭高斯,闭式解 * ELBO —— 重建项对 q 可解析求期望,KL 两项都是高斯闭式 有解析解才能验证「差的那一项到底是什么」,靠采样是验不出来的。 只依赖 numpy,直接 `python elbo_identity.py` 即可运行。 """ import numpy as np RNG = np.random.default_rng(20260925) # ── 模型:z ~ N(0, I_K),x|z ~ N(W z + b, sigma^2 I_D) ── D, K = 6, 3 SIGMA = 0.35 # 解码器观测噪声标准差 SIGMA2 = SIGMA ** 2 W = RNG.normal(0.0, 1.0, size=(D, K)) * 0.9 # 解码器权重 b = RNG.normal(0.0, 0.3, size=(D,)) # 解码器偏置 def make_one_x(rng=None): """从真实的边际分布里采一个 x,并返回它的解析 log p(x)。 rng 可显式传入:这样 reparam_gradients / make_figures 拿到的是同一个 x, 不会因为模块级 RNG 被别人消耗过而对不上数字。 """ if rng is None: rng = RNG z = rng.normal(size=(K,)) x = W @ z + b + SIGMA * rng.normal(size=(D,)) cov = W @ W.T + SIGMA2 * np.eye(D) sign, logdet = np.linalg.slogdet(cov) d = x - b quad = d @ np.linalg.solve(cov, d) log_px = -0.5 * (D * np.log(2 * np.pi) + logdet + quad) return x, log_px def true_posterior(x): """p(z|x) 的闭式解:precision = I + W^T W / sigma^2。""" prec = np.eye(K) + W.T @ W / SIGMA2 cov = np.linalg.inv(prec) mean = cov @ (W.T @ (x - b) / SIGMA2) return mean, cov # ── 变分后验 q(z|x) = N(m, diag(s^2)):故意用「对角」错误指定 ── def make_q(x, rng=None): """取真实后验的均值,协方差只留对角线并加一点扰动。 真实后验是有相关性的(cov 非对角),q 强行对角 => KL(q||p(z|x)) > 0, 这正是我们要留出来的那条缝。 """ if rng is None: rng = RNG mean, cov = true_posterior(x) m = mean + 0.15 * rng.normal(size=(K,)) # 均值也偏一点 var = np.diag(cov).copy() * 0.8 + 0.05 # 方差偏小 => q 过自信 return m, np.sqrt(var) def log_gauss_diag(z, mean, var, ): """log N(z; mean, diag(var))(省略常数也行,但这里算全)。""" return float(-0.5 * np.sum(np.log(2 * np.pi * var) + (z - mean) ** 2 / var)) def log_joint(x, z): """log p(x, z) = log p(x|z) + log p(z)。""" r = x - (W @ z + b) log_lik = -0.5 * (D * np.log(2 * np.pi * SIGMA2) + r @ r / SIGMA2) log_prior = -0.5 * (K * np.log(2 * np.pi) + z @ z) return log_lik + log_prior def elbo_analytic(x, m, s): """ELBO 的解析值。 E_q[log p(x|z)] 里 ||x - Wz - b||^2 对 q 求期望: ||x - Wm - b||^2 + tr(W^T W diag(s^2)) 第二项是采样噪声经过解码器放大出来的那部分,容易被漏掉。 """ var = s ** 2 resid = x - (W @ m + b) recon = -0.5 * (D * np.log(2 * np.pi * SIGMA2) + (resid @ resid + np.trace(W.T @ W @ np.diag(var))) / SIGMA2) kl = 0.5 * np.sum(var + m ** 2 - 1.0 - np.log(var)) return float(recon - kl), float(recon), float(kl) def kl_gauss_gauss(m, s, mean2, cov2): """KL(N(m, diag(s^2)) || N(mean2, cov2)),两个高斯的闭式。""" var = s ** 2 diff = mean2 - m cov2_inv = np.linalg.inv(cov2) trace = float(np.trace(cov2_inv @ np.diag(var))) quad = float(diff @ cov2_inv @ diff) _, logdet2 = np.linalg.slogdet(cov2) return 0.5 * (trace + quad - K + logdet2 - np.sum(np.log(var))) def mc_elbo(x, m, s, L, rng): """蒙特卡洛 ELBO:L 个样本取平均(L=1 就是训练时实际用的那个)。""" eps = rng.normal(size=(L, K)) z = m + s * eps log_lik = np.array([-0.5 * (D * np.log(2 * np.pi * SIGMA2) + np.sum((x - (W @ zz + b)) ** 2) / SIGMA2) for zz in z]) var = s ** 2 kl = 0.5 * np.sum(var + m ** 2 - 1.0 - np.log(var)) return float(log_lik.mean() - kl) def iwae_log_bound(x, m, s, k, rng): """IWAE 的 k 样本界:log (1/k) sum_l p(x,z_l)/q(z_l),用 logsumexp 稳算。""" eps = rng.normal(size=(k, K)) z = m + s * eps var = s ** 2 logw = np.array([log_joint(x, zz) - log_gauss_diag(zz, m, var) for zz in z]) mx = logw.max() return float(mx + np.log(np.mean(np.exp(logw - mx)))) # ── 下面两个表被 main 与 make_figures 共用,种子固定 => 数字永远一致 ── MC_LS = (1, 4, 16, 64) MC_REP = 4000 IW_KS = (1, 5, 50, 500) IW_REP = 6000 STATE_SEED = 20260925 def state(): """文章与配图共用的那一份状态:同一个 x、同一个 q。 用固定种子重放,保证单独跑本脚本与跑 make_figures 拿到同一组数字。 """ rng = np.random.default_rng(STATE_SEED) x, log_px = make_one_x(rng) m, s = make_q(x, rng) post_mean, post_cov = true_posterior(x) elbo, recon, kl = elbo_analytic(x, m, s) gap = kl_gauss_gauss(m, s, post_mean, post_cov) return dict(x=x, log_px=log_px, m=m, s=s, elbo=elbo, recon=recon, kl=kl, gap=gap) def mc_table(st): """蒙特卡洛 ELBO 的均值 / 标准差 / 越过 log p(x) 的比例。""" rng = np.random.default_rng(7) out = {} for L in MC_LS: vals = np.array([mc_elbo(st["x"], st["m"], st["s"], L, rng) for _ in range(MC_REP)]) out[L] = dict(mean=float(vals.mean()), std=float(vals.std()), frac=float(np.mean(vals > st["log_px"])), vals=vals) return out def iwae_table(st): """IWAE 的 k 样本界。""" rng = np.random.default_rng(13) out = {} for k in IW_KS: vals = np.array([iwae_log_bound(st["x"], st["m"], st["s"], k, rng) for _ in range(IW_REP)]) out[k] = dict(mean=float(vals.mean()), std=float(vals.std()), gap=float(st["log_px"] - vals.mean()), vals=vals) return out def main(): print("=" * 68) print("ELBO 恒等式核对 log p(x) = ELBO + KL(q || p(z|x))") print("=" * 68) print(f"模型: D={D}, K={K}, sigma={SIGMA}") st = state() x, log_px, m, s = st["x"], st["log_px"], st["m"], st["s"] elbo, recon, kl, gap = st["elbo"], st["recon"], st["kl"], st["gap"] print(f"\n[1] 解析三项") print(f" E_q[log p(x|z)] = {recon: .6f}") print(f" KL(q || p(z)) = {kl: .6f}") print(f" ELBO = {elbo: .6f}") print(f" KL(q || p(z|x)) = {gap: .6f}") print(f" ELBO + KL(q||p(z|x)) = {elbo + gap: .6f}") print(f" log p(x) = {log_px: .6f}") print(f" |误差| = {abs(elbo + gap - log_px): .3e}") # ── 蒙特卡洛波动:ELBO 的估计量是无偏的,不是恒小于 log p(x) ── print(f"\n[2] 蒙特卡洛估计的波动(每项 {MC_REP} 次重复)") print(f" {'L':>5s} {'mean':>12s} {'std':>10s} {'超过 log p(x) 的比例':>22s}") mc = mc_table(st) for L in MC_LS: r = mc[L] print(f" {L:5d} {r['mean']:12.5f} {r['std']:10.5f} {r['frac']:21.1%}") print(f" log p(x) = {log_px:.5f}(上界本身),ELBO(解析) = {elbo:.5f}") # ── IWAE:k 越大越紧 ── print(f"\n[3] IWAE 的 k 样本界(每项 {IW_REP} 次重复)") print(f" {'k':>5s} {'mean L_k':>12s} {'std':>9s} {'log p(x) - L_k':>16s}") iw = iwae_table(st) for k in IW_KS: r = iw[k] print(f" {k:5d} {r['mean']:12.5f} {r['std']:9.5f} {r['gap']:16.5f}") print(f" 参考:单样本 MC ELBO 的均值 = {elbo:.5f}(= ELBO 解析值,不随 k 变)") print("\n结论:k 增大时 L_k 单调逼近 log p(x),但永远不越过;" "而「取平均再 log」和「log 再取平均」是两件事。") if __name__ == "__main__": main() reparam_gradients.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """重参数化 vs 得分函数(REINFORCE)梯度估计量,方差实测。 要估的量:∇_φ E_{q_φ}[log p(x|z)],φ = (m, ℓ),ℓ = log s。 这一项没法解析求的时候只能采样,两种采样方式的方差差好几个数量级。 真实梯度在这个高斯线性模型里能解析算出来,所以「谁对谁错」有标准答案, 剩下的差别就纯粹是方差。 两部分实验: A. 固定一个小模型(D=6),看三种估计量的分布 B. 扫数据维度 D,看方差比怎么长——这里才是重参数化真正救命的地方 只依赖 numpy,直接 `python reparam_gradients.py` 即可运行。 """ import numpy as np import elbo_identity as EI from elbo_identity import ( D, K, SIGMA2, W, b, elbo_analytic, make_one_x, make_q, ) N_TRIAL = 40000 # ══════════════════════════════════════════════════════════════ # A. 小模型:三种估计量的分布 # ══════════════════════════════════════════════════════════════ def true_grad(x, m, s, Wm, bm, sigma2): """∇_m 与 ∇_ℓ 的解析梯度,ℓ = log s。""" g_m = Wm.T @ (x - (Wm @ m + bm)) / sigma2 g_ell = -(s ** 2) * np.diag(Wm.T @ Wm) / sigma2 return g_m, g_ell def log_lik_rows(x, z, Wm, bm, Dm, sigma2): """批量算 log p(x|z),z: [n, K]。""" r = x - (z @ Wm.T + bm) return -0.5 * (Dm * np.log(2 * np.pi * sigma2) + np.einsum("nd,nd->n", r, r) / sigma2) def sample_reparam(x, m, s, Wm, bm, sigma2, rng, n): """重参数化:z = m + s ⊙ ε,梯度经 z 反传。""" eps = rng.normal(size=(n, m.size)) z = m + s * eps score_z = (x - (z @ Wm.T + bm)) @ Wm / sigma2 # ∇_z log p(x|z) return score_z, score_z * s * eps # ∂z/∂m=1, ∂z/∂ℓ=s⊙ε def sample_score(x, m, s, Wm, bm, Dm, sigma2, rng, n, baseline=None): """得分函数:∇ E[f] = E[f · ∇ log q],f = log p(x|z)。""" eps = rng.normal(size=(n, m.size)) z = m + s * eps f = log_lik_rows(x, z, Wm, bm, Dm, sigma2) if baseline is not None: f = f - baseline return f[:, None] * (eps / s), f[:, None] * (eps ** 2 - 1.0) def summarize(name, gm_hat, ge_hat, g_m, g_ell): var = float(np.var(gm_hat, axis=0).sum() + np.var(ge_hat, axis=0).sum()) gnorm = float(np.linalg.norm(np.concatenate([g_m, g_ell]))) err = float(np.linalg.norm(np.concatenate([gm_hat.mean(0) - g_m, ge_hat.mean(0) - g_ell]))) / gnorm n_needed = int(np.ceil((var ** 0.5 / (0.05 * gnorm)) ** 2)) print(f" {name:<24s} 方差和={var:11.2f} 均值相对误差={err:6.2%} " f"达到5%需 n≈{n_needed}") return var def part_a(): print("=" * 68) print("A. 小模型(D=6, σ=0.35):三种估计量的分布") print("=" * 68) # 与 elbo_identity.state() 同一颗种子 => 同一个 x、同一个 q rng_state = np.random.default_rng(EI.STATE_SEED) x, _ = make_one_x(rng_state) m, s = make_q(x, rng_state) g_m, g_ell = true_grad(x, m, s, W, b, SIGMA2) _, recon, kl = elbo_analytic(x, m, s) print(f"\n真实梯度(解析)") print(f" ∇_m = {np.array2string(g_m, precision=4)}") print(f" ∇_ℓ = {np.array2string(g_ell, precision=4)}") print(f" E_q[log p(x|z)] = {recon:.4f},KL(q||p(z)) = {kl:.4f}") rng = np.random.default_rng(11) n = N_TRIAL print(f"\n估计量对比({n} 次独立采样,方差 = 6 个参数分量方差之和)") gm_r, ge_r = sample_reparam(x, m, s, W, b, SIGMA2, rng, n) gm_s, ge_s = sample_score(x, m, s, W, b, D, SIGMA2, rng, n) gm_sb, ge_sb = sample_score(x, m, s, W, b, D, SIGMA2, rng, n, baseline=recon) v_r = summarize("重参数化", gm_r, ge_r, g_m, g_ell) v_s = summarize("得分函数(无 baseline)", gm_s, ge_s, g_m, g_ell) v_sb = summarize("得分函数(预言机 baseline)", gm_sb, ge_sb, g_m, g_ell) print(f"\n 方差比:得分函数 / 重参数化 = {v_s / v_r:.1f}×") print(f" 加 baseline 之后 = {v_sb / v_r:.1f}×") print(f"\n∇_m 第 0 个分量(真值 {g_m[0]:.4f})的估计分布") for name, g in (("重参数化", gm_r[:, 0]), ("得分函数", gm_s[:, 0]), ("得分函数+baseline", gm_sb[:, 0])): print(f" {name:<20s} mean={g.mean():9.4f} std={g.std():8.3f} " f"min={g.min():9.2f} max={g.max():9.2f}") print(f"\n均值误差随样本数的收敛(∇_m 第 0 个分量)") print(f" {'n':>7s} {'重参数化':>12s} {'得分函数':>12s}") for n2 in (10, 100, 1000, 10000, N_TRIAL): print(f" {n2:7d} {abs(gm_r[:n2, 0].mean() - g_m[0]):12.5f} " f"{abs(gm_s[:n2, 0].mean() - g_m[0]):12.5f}") return dict(g_m0=g_m[0], gm_r=gm_r[:, 0], gm_s=gm_s[:, 0], gm_sb=gm_sb[:, 0], var_r=v_r, var_s=v_s, var_sb=v_sb) # ══════════════════════════════════════════════════════════════ # B. 方差比随数据维度 D 怎么长 # ══════════════════════════════════════════════════════════════ def build_model(Dm, Km, seed): """构造一个 decoder 列范数归一化的线性高斯模型,让梯度尺度不随 D 漂。""" rng = np.random.default_rng(seed) Wm = rng.normal(size=(Dm, Km)) Wm /= np.linalg.norm(Wm, axis=0, keepdims=True) # 每列范数 = 1 bm = rng.normal(0.0, 0.3, size=(Dm,)) return Wm, bm def part_b(): print("\n" + "=" * 68) print("B. 方差比随数据维度 D 的变化(σ²=1,decoder 列范数归一化)") print("=" * 68) rng = np.random.default_rng(23) rows = [] print(f"\n {'D':>6s} {'重参数化':>12s} {'得分函数':>14s} {'比值':>10s} " f"{'+baseline 方差':>14s} {'比值':>10s}") for Dm in (8, 64, 512, 2048): Wm, bm = build_model(Dm, K, seed=1000 + Dm) z_true = rng.normal(size=(K,)) x = Wm @ z_true + bm + rng.normal(size=(Dm,)) # σ² = 1 m = Wm.T @ (x - bm) # 一个合理的变分均值 s = np.full((K,), 0.6) sigma2 = 1.0 n = 4000 gm_r, ge_r = sample_reparam(x, m, s, Wm, bm, sigma2, rng, n) gm_s, ge_s = sample_score(x, m, s, Wm, bm, Dm, sigma2, rng, n) # 预言机 baseline:E_q[log p(x|z)] 的解析值 resid = x - (Wm @ m + bm) recon = -0.5 * (Dm * np.log(2 * np.pi * sigma2) + (resid @ resid + np.trace(Wm.T @ Wm @ np.diag(s ** 2))) / sigma2) gm_sb, ge_sb = sample_score(x, m, s, Wm, bm, Dm, sigma2, rng, n, baseline=recon) v_r = float(np.var(gm_r, axis=0).sum() + np.var(ge_r, axis=0).sum()) v_s = float(np.var(gm_s, axis=0).sum() + np.var(ge_s, axis=0).sum()) v_sb = float(np.var(gm_sb, axis=0).sum() + np.var(ge_sb, axis=0).sum()) rows.append((Dm, v_r, v_s, v_sb)) print(f" {Dm:6d} {v_r:12.2f} {v_s:14.2f} {v_s / v_r:9.1f}× " f"{v_sb:14.2f} {v_sb / v_r:9.1f}×") print("\n 为什么长这么快:f = log p(x|z) 的均值里有一大坨与参数无关的") print(" 「底噪」(-D/2·log 2πσ² 占了主要部分),它进到 f·∇log q 里按平方") print(" 放大;∇_z f 对它求导恒为 0,所以重参数化天然免疫。baseline 减掉的") print(" 也正是这一坨——减完比值只剩 5 倍左右,说明差距主要来自常数项,") print(" 不是来自「采样本身」。") return rows if __name__ == "__main__": a = part_a() rows = part_b() beta_kl_weight.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """β-VAE 的 KL 加权:一个能解到最优的玩具模型,看潜变量维度怎么死掉。 构造:数据每个坐标独立,x_j ~ N(0, λ_j),潜变量每个坐标也独立, q(z_j | x_j) = N(a_j x_j, s_j²),p(x_j | z_j) = N(b_j z_j, σ²) 这样每个维度就是一份独立副本,最优解可以逐个维度单独求,不用训练神经网络。 单维目标(要最大化),把 λ 与 σ² 都当作常数: L(a, b, ℓ) = -[λ(1 - b a)² + b² s²] / (2σ²) - β · KL KL = 0.5 · (a² λ + s² - 1 - log s²), s = exp(ℓ) 梯度全部解析,脚本里再用有限差分核对一遍,避免推错。 只依赖 numpy,直接 `python beta_kl_weight.py` 即可运行。 """ import numpy as np SIGMA2 = 0.25 # 解码器观测噪声方差 def kl_of(a, lam, s): """KL(N(a x, s²) || N(0, 1)),对 x ~ N(0, λ) 取期望后的形式。""" return 0.5 * (a ** 2 * lam + s ** 2 - 1.0 - np.log(s ** 2)) def recon_mse(a, b, lam, s): """E_{x,z}[(x - b z)²],重建误差的期望。""" return lam * (1.0 - b * a) ** 2 + b ** 2 * s ** 2 def objective(p, lam, beta, free_bits=None): """目标函数值(越大越好)。free_bits=C 时 KL 项取 max(KL, C)。""" a, b, ell = p s = np.exp(ell) kl = kl_of(a, lam, s) kl_eff = max(kl, free_bits) if free_bits is not None else kl return -(recon_mse(a, b, lam, s)) / (2 * SIGMA2) - beta * kl_eff def grad(p, lam, beta, free_bits=None): """解析梯度 ∇(∂L/∂a, ∂L/∂b, ∂L/∂ℓ)。""" a, b, ell = p s2 = np.exp(2 * ell) g = np.zeros(3) g[0] = lam * b * (1.0 - b * a) / SIGMA2 g[1] = (lam * a * (1.0 - b * a) - b * s2) / SIGMA2 g[2] = -b ** 2 * s2 / SIGMA2 if free_bits is None or kl_of(a, lam, np.exp(ell)) > free_bits: # KL 项对 (a, b, ℓ) 的梯度 g[0] -= beta * a * lam g[2] -= beta * (s2 - 1.0) return g def grad_fd(p, lam, beta, free_bits=None, h=1e-6): """中心差分梯度,用来核对解析梯度有没有推错。""" g = np.zeros(3) for i in range(3): e = np.zeros(3) e[i] = h g[i] = (objective(p + e, lam, beta, free_bits) - objective(p - e, lam, beta, free_bits)) / (2 * h) return g def optimize(lam, beta, free_bits=None, n_init=9, n_iter=12000, seed=0): """多起点梯度上升,返回最好的 (a, b, s)。 必须多起点:(a, b) = (0, 0) 是「维度死亡」解,单起点容易掉进去出不来。 """ rng = np.random.default_rng(seed) best_p, best_v = None, -np.inf for k in range(n_init): if k == 0: p = np.array([0.9, 0.9, np.log(0.6)]) elif k == 1: p = np.array([0.0, 0.0, 0.0]) # 死亡解,也让它试试 else: p = np.array([rng.uniform(-1.5, 1.5), rng.uniform(-1.5, 1.5), rng.uniform(-1.2, 0.5)]) lr = 0.02 for t in range(n_iter): g = grad(p, lam, beta, free_bits) g = np.clip(g, -50.0, 50.0) # 梯度裁剪:b²s²/σ² 那一项能把参数炸飞 p = p + lr * g # 数值保护,别让 exp(ℓ) 或 (a, b) 爆掉 p[0] = float(np.clip(p[0], -8.0, 8.0)) p[1] = float(np.clip(p[1], -8.0, 8.0)) p[2] = float(np.clip(p[2], -4.0, 1.0)) if t % 1000 == 999: lr *= 0.6 # 收尾再磨一遍:β 很小的时候收敛慢,不磨的话 s² 会差到 1e-3 lr = 1e-4 for t in range(6000): g = np.clip(grad(p, lam, beta, free_bits), -50.0, 50.0) p = p + lr * g v = objective(p, lam, beta, free_bits) if v > best_v: best_v, best_p = v, p a, b, ell = best_p return float(a), float(b), float(np.exp(ell)), float(best_v) def active_unit(a, lam, thr=0.01): """Higgins 的 active unit 判据:Cov_x(E_q[z]) = a²λ 是否超过阈值。""" return a ** 2 * lam > thr # 文章与配图共用的扫描范围 LAMS = [4.0, 1.0, 0.25, 0.0625] BETAS = [0.25, 0.5, 1.0, 2.0, 3.0, 4.0, 6.0, 8.0, 12.0] def sweep(lams, betas, verbose=True): """扫 (λ, β) 网格,返回 {(λ, β): (a, b, s, MSE, KL, a²λ)}。""" table = {} if verbose: print("\n[1] β 扫描:每个 (λ, β) 的最优解") print(f"\n {'λ':>7s} {'β':>5s} {'a':>8s} {'b':>8s} {'s':>8s} " f"{'重建MSE':>9s} {'KL':>8s} {'a²λ':>8s} {'存活':>5s}") for lam in lams: for beta in betas: a, b, s, _ = optimize(lam, beta, seed=17) kl = kl_of(a, lam, s) mse = recon_mse(a, b, lam, s) au = a ** 2 * lam table[(lam, beta)] = (a, b, s, mse, kl, au) if verbose: print(f" {lam:7.4f} {beta:5.2f} {a:8.4f} {b:8.4f} {s:8.4f} " f"{mse:9.4f} {kl:8.4f} {au:8.4f} " f"{'是' if active_unit(a, lam) else '死':>5s}") return table def main(): print("=" * 72) print("β-VAE 的 KL 加权:潜变量维度怎么死掉") print(f"σ² = {SIGMA2}(解码器噪声),单维独立副本,多起点梯度上升求最优") print("=" * 72) # ── 先核对解析梯度 ── print("\n[0] 解析梯度 vs 有限差分") for lam, beta in ((4.0, 1.0), (0.25, 4.0)): p = np.array([0.7, -0.4, np.log(0.8)]) ga, gf = grad(p, lam, beta), grad_fd(p, lam, beta) print(f" λ={lam:<5} β={beta:<4} 解析={np.array2string(ga, precision=5)} " f"差分={np.array2string(gf, precision=5)} 最大偏差={np.abs(ga - gf).max():.2e}") table = sweep(LAMS, BETAS) print("\n[2] 每个 λ 的「死亡阈值」:β 到多大时这个维度不再被用") for lam in LAMS: dead = [beta for beta in BETAS if not active_unit(table[(lam, beta)][0], lam)] if dead: print(f" λ={lam:<7.4f} 从 β={min(dead)} 起死亡" f"(β=1 时 KL={table[(lam, 1.0)][4]:.4f})") else: print(f" λ={lam:<7.4f} 在 β≤{max(BETAS)} 内都存活") # ── 闭式解核对:存活时 s² = βσ²/λ,a²λ = 1 - s²,重建 MSE = βσ² ── print("\n[2b] 闭式解核对(存活的格子才成立)") print(f" {'λ':>7s} {'β':>5s} {'s²实测':>9s} {'βσ²/λ':>9s} " f"{'a²λ实测':>9s} {'1-s²':>9s} {'MSE实测':>9s} {'βσ²':>9s}") worst = 0.0 for lam in LAMS: for beta in BETAS: a, b, s, _, kl, au = table[(lam, beta)] if not active_unit(a, lam): continue s2_pred = beta * SIGMA2 / lam rows = (s ** 2, s2_pred, au, 1 - s ** 2, recon_mse(a, b, lam, s), beta * SIGMA2) print(f" {lam:7.4f} {beta:5.2f} " + " ".join(f"{v:9.5f}" for v in rows)) worst = max(worst, abs(s ** 2 - s2_pred), abs(au - (1 - s ** 2)), abs(recon_mse(a, b, lam, s) - beta * SIGMA2)) print(f" 最大偏差 = {worst:.2e} → 存活时最优解确实落在闭式上") print(f" 推论:维度存活条件 s² < 1 ⟺ β < λ/σ²,即该维的信噪比要盖过 KL 的价") for lam in LAMS: print(f" λ={lam:<7.4f} 理论死亡阈值 β* = λ/σ² = {lam / SIGMA2:.2f}") # ── free bits:把 KL 压在下界,维度就不会死 ── print("\n[3] free bits 对照(β=4,KL 项取 max(KL, C))") print(f" {'λ':>7s} {'C=0(纯 β)':>22s} {'C=0.05':>22s} {'C=0.2':>22s}") for lam in LAMS: cells = [] for C in (None, 0.05, 0.2): a, b, s, _ = optimize(lam, 4.0, free_bits=C, seed=17) cells.append(f"a={a:5.3f} KL={kl_of(a, lam, s):5.3f} " f"MSE={recon_mse(a, b, lam, s):5.3f}") print(f" {lam:7.4f} " + " ".join(f"{c:>22s}" for c in cells)) print("\n[4] 一个直觉:β 大了以后,模型宁可不要这个维度") print(" λ=0.0625 时这个坐标的信息量本来就小于解码器噪声 σ²=0.25,") print(" 用它换来的重建收益抵不上 KL 的代价,最优解就是把它整条关掉。") if __name__ == "__main__": main() make_figures.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """生成「变分下界与重参数化」的三张解释图。 数值全部来自同目录的三个脚本(elbo_identity / reparam_gradients / beta_kl_weight),这里只负责把它们画出来——改了那几个脚本这里要重跑, 避免图与正文数字不一致。 三张图分别回答: 1. ELBO 离真实目标差多少,差的那一项是什么,多采样能不能补上 2. 重参数化到底省了多少方差,以及这个差距随数据维度怎么长 3. KL 项加权的 β 是怎么把潜变量维度一刀一刀切掉的 只依赖 numpy + matplotlib,直接 `python make_figures.py` 即可运行。 """ from pathlib import Path import matplotlib.pyplot as plt import numpy as np from matplotlib.patches import Rectangle import beta_kl_weight as BK import elbo_identity as EI import reparam_gradients as RG ROOT = Path(__file__).resolve().parents[1] OUT = ROOT / "figures" OUT.mkdir(exist_ok=True) plt.rcParams.update({ "font.sans-serif": ["PingFang SC", "Arial Unicode MS", "DejaVu Sans"], "axes.unicode_minus": False, "figure.dpi": 160, }) INK = "#1f2937" C_ELBO = "#2f6fb0" # ELBO:蓝 C_GAP = "#e0a03c" # 缺口:橙 C_BAD = "#d1495b" # 越界 / 死亡:红 C_OK = "#2f9e6f" # 收紧 / 存活:绿 C_SCORE = "#8b5cf6" # 得分函数:紫 def style(ax, title, xlabel=None, ylabel=None): ax.set_title(title, fontsize=13.0, weight="bold", color=INK, pad=10) if xlabel: ax.set_xlabel(xlabel, fontsize=11, color="#475569") if ylabel: ax.set_ylabel(ylabel, fontsize=11, color="#475569") ax.tick_params(labelsize=10, colors="#64748b") for s in ("top", "right"): ax.spines[s].set_visible(False) for s in ("left", "bottom"): ax.spines[s].set_color("#cbd5e1") ax.grid(axis="y", color="#eef2f7", lw=1.0) ax.set_axisbelow(True) # ══════════════════════════════════════════════════════════════ # 图 1:ELBO 离 log p(x) 差多少 # ══════════════════════════════════════════════════════════════ def fig_elbo_gap(): st = EI.state() mc = EI.mc_table(st) iw = EI.iwae_table(st) log_px, elbo, gap, kl_prior = st["log_px"], st["elbo"], st["gap"], st["kl"] fig, axes = plt.subplots(1, 3, figsize=(16.2, 4.9)) fig.patch.set_facecolor("white") # ── (a) 分解:两条横向长条 ── ax = axes[0] style(ax, "(a) log p(x) 拆成两截", "nats(对数似然,0 在右边)") rows = [ ("ELBO(能算,优化它)", elbo, C_ELBO, "white"), ("log p(x)(真想要的,算不出来)", log_px, C_BAD, "white"), ("KL(q‖p(z)):loss 里那一项", -kl_prior, "#94a3b8", "white"), ("KL(q‖p(z|x)):上面两条的差", -gap, C_GAP, "white"), ] for i, (name, v, color, tc) in enumerate(rows): y = len(rows) - 1 - i ax.barh(y, -v, left=v, height=0.52, color=color) ax.text(v + 0.18, y, f"{abs(v):.3f}", ha="left", va="center", fontsize=10, color=INK, weight="bold") ax.axvline(log_px, color=C_BAD, ls="--", lw=1.3) ax.axvline(elbo, color=C_ELBO, ls="--", lw=1.3) ax.annotate("", xy=(elbo, 3.28), xytext=(log_px, 3.28), arrowprops=dict(arrowstyle="<->", color="#a16207", lw=1.8)) ax.text((elbo + log_px) / 2, 3.75, "差 = 1.985", ha="center", va="bottom", fontsize=10.5, color="#a16207", weight="bold", bbox=dict(fc="white", ec="none", alpha=0.9, pad=1.0)) ax.text(-9.35, 0.68, "上面两条之差 = 下面橙色那条,\n不是灰色那条——这两个 KL 常被混为一谈", ha="left", va="center", fontsize=9, color="#64748b", linespacing=1.5) ax.set_yticks(range(len(rows))) ax.set_yticklabels([r[0] for r in rows][::-1], fontsize=9.5) ax.set_xlim(-9.6, 1.5) ax.set_ylim(-0.6, 3.9) ax.grid(axis="x", color="#eef2f7", lw=1.0) ax.grid(axis="y", visible=False) # ── (b) 单样本 MC 会越过上界 ── ax = axes[1] style(ax, f"(b) L=1 的 MC 估计({EI.MC_REP} 次)", "单次估计值", "频数") vals = mc[1]["vals"] ax.hist(vals, bins=70, color="#cbd5e1", edgecolor="white", lw=0.4) over = vals[vals > log_px] ax.hist(over, bins=70, color=C_BAD, alpha=0.85, label=f"{mc[1]['frac']:.1%} 越过了 log p(x)") ax.axvline(elbo, color=C_ELBO, lw=2.0, label=f"ELBO = {elbo:.3f}") ax.axvline(log_px, color=C_BAD, lw=2.0, ls="--", label=f"log p(x) = {log_px:.3f}") ax.set_xlim(-26, 7) ax.legend(fontsize=9, frameon=False, loc="upper left") ax.set_yscale("log") ax.set_ylim(0.7, 900) ax.text(0.03, 0.05, "估计量是无偏的:它在 ELBO 周围晃,\n不是「恒小于 log p(x)」", transform=ax.transAxes, ha="left", va="bottom", fontsize=9.5, color="#64748b", linespacing=1.5) # ── (c) IWAE 收紧 ── ax = axes[2] style(ax, f"(c) 多采样收紧(每项 {EI.IW_REP} 次)", "k(每次采几个 z)", "nats") ks = np.array(EI.IW_KS, dtype=float) means = np.array([iw[k]["mean"] for k in EI.IW_KS]) stds = np.array([iw[k]["std"] for k in EI.IW_KS]) ax.errorbar(ks, means, yerr=stds, fmt="o-", color=C_OK, lw=2.0, capsize=4, markersize=6) ax.axhline(log_px, color=C_BAD, ls="--", lw=1.6, label=f"log p(x) = {log_px:.3f}") ax.axhline(elbo, color=C_ELBO, ls=":", lw=1.6, label=f"ELBO = {elbo:.3f}") ax.set_xscale("log") ax.set_ylim(-9.4, -4.7) ax.set_xlim(0.7, 900) for k, m in zip(ks, means): ax.annotate(f"还差 {log_px - m:.3f}", (k, m), textcoords="offset points", xytext=(8, 12 if k > 1 else -18), fontsize=9, color="#64748b", bbox=dict(fc="white", ec="none", alpha=0.85, pad=0.6)) ax.legend(fontsize=9, frameon=False, loc="lower right") ax.text(0.03, 0.06, "k 越大越紧,但永远不越过虚线", transform=ax.transAxes, fontsize=9.5, color="#64748b") fig.suptitle("这张图要看什么:ELBO 与真实目标之间那道缝," "一半来自 q 的表达力,一半来自「只采一个样本」", fontsize=11.5, color="#64748b", y=1.02) fig.tight_layout() p = OUT / "elbo_gap.png" fig.savefig(p, bbox_inches="tight", facecolor="white") plt.close(fig) print(f" ✓ {p.name}") # ══════════════════════════════════════════════════════════════ # 图 2:重参数化省了多少方差 # ══════════════════════════════════════════════════════════════ def fig_grad_variance(): a = RG.part_a() rows = RG.part_b() fig, axes = plt.subplots(1, 2, figsize=(13.4, 5.0)) fig.patch.set_facecolor("white") # ── (a) 估计分布 ── ax = axes[0] style(ax, "(a) ∂L/∂m 第 0 个分量的估计分布", "估计值", "频数") bins = np.linspace(-70, 70, 90) for vals, color, name in ((a["gm_r"], C_ELBO, "重参数化"), (a["gm_s"], C_SCORE, "得分函数"), (a["gm_sb"], C_GAP, "得分函数+baseline")): ax.hist(vals, bins=bins, histtype="step", lw=1.8, color=color, label=f"{name} std={vals.std():.1f}") ax.axvline(a["g_m0"], color=C_OK, lw=2.2, label=f"真值 = {a['g_m0']:.3f}") ax.set_xlim(-70, 70) ax.set_yscale("log") ax.legend(fontsize=9.5, frameon=False, loc="upper left") ax.text(0.98, 0.05, "三条曲线均值都对,差的是腰围:\n" f"std = {a['gm_r'].std():.1f} / {a['gm_s'].std():.1f} / " f"{a['gm_sb'].std():.1f}", transform=ax.transAxes, ha="right", va="bottom", fontsize=9.5, color="#64748b", linespacing=1.5) # ── (b) 随维度怎么长 ── ax = axes[1] style(ax, "(b) 方差随数据维度 D 怎么长", "D(数据维度)", "梯度方差(6 个分量之和)") Ds = np.array([r[0] for r in rows], dtype=float) vr = np.array([r[1] for r in rows]) vs = np.array([r[2] for r in rows]) vsb = np.array([r[3] for r in rows]) ax.loglog(Ds, vs, "o-", color=C_SCORE, lw=2.0, markersize=7, label="得分函数(无 baseline)") ax.loglog(Ds, vsb, "s--", color=C_GAP, lw=2.0, markersize=6, label="得分函数 + baseline") ax.loglog(Ds, vr, "o-", color=C_ELBO, lw=2.0, markersize=7, label="重参数化") # 斜率 2 的参考线 ref = vs[0] * (Ds / Ds[0]) ** 2 ax.loglog(Ds, ref, ":", color="#cbd5e1", lw=2.0, label="∝ D² 参考线") for d, v, r in zip(Ds, vs, vr): ratio = v / r txt = f"{ratio:,.0f}×" if ratio < 1e4 else f"{ratio:.1e}×" ax.annotate(txt, (d, v), textcoords="offset points", xytext=(8, -14), fontsize=9.5, color=C_SCORE, weight="bold", bbox=dict(fc="white", ec="none", alpha=0.85, pad=0.6)) ax.set_ylim(2e-1, 1e10) ax.legend(fontsize=9.5, frameon=False, loc="upper left") ax.text(0.97, 0.03, "重参数化几乎与 D 无关", transform=ax.transAxes, fontsize=9.5, color=C_ELBO, ha="right", va="bottom", weight="bold") fig.suptitle("这张图要看什么:重参数化不是「更准」,是「腰围更小」——" "同样一次采样,梯度离真值近一个数量级", fontsize=11.5, color="#64748b", y=1.02) fig.tight_layout() p = OUT / "grad_variance.png" fig.savefig(p, bbox_inches="tight", facecolor="white") plt.close(fig) print(f" ✓ {p.name}") # ══════════════════════════════════════════════════════════════ # 图 3:β 怎么把潜变量维度切掉 # ══════════════════════════════════════════════════════════════ def fig_beta_tradeoff(): table = BK.sweep(BK.LAMS, BK.BETAS, verbose=False) s2 = BK.SIGMA2 fig, axes = plt.subplots(1, 2, figsize=(13.4, 5.0)) fig.patch.set_facecolor("white") # ── (a) 重建误差 vs β ── ax = axes[0] style(ax, "(a) β 越大,重建越差——直到某一维被直接放弃", "β(KL 项的权重)", "该维的重建 MSE") colors = ["#2f6fb0", "#2f9e6f", "#e0a03c", "#d1495b"] bgrid = np.linspace(0.2, 14, 100) ax.plot(bgrid, bgrid * s2, ":", color="#94a3b8", lw=1.8, label="理论:MSE = βσ²(还活着)") for lam, c in zip(BK.LAMS, colors): betas, mses, alive = [], [], [] for beta in BK.BETAS: a, b, s, mse, kl, au = table[(lam, beta)] betas.append(beta) mses.append(mse) alive.append(BK.active_unit(a, lam)) betas = np.array(betas) mses = np.array(mses) alive = np.array(alive) ax.plot(betas[alive], mses[alive], "o-", color=c, lw=2.0, markersize=6, label=f"λ={lam}") if (~alive).any(): ax.plot(betas[~alive], mses[~alive], "s", color=c, markersize=7, mfc="white", mew=1.8) ax.axhline(mses[~alive][0], color=c, ls="--", lw=1.0, alpha=0.45) ax.set_xscale("log") ax.set_yscale("log") ax.set_ylim(0.028, 7.0) ax.legend(fontsize=9, frameon=False, loc="upper left", ncol=2) ax.text(0.5, 0.012, "实心圆 = 该维还活着(MSE 贴着 βσ² 往上走);" "空心方块 = 该维已死,MSE 停在 λ", transform=ax.transAxes, ha="center", va="bottom", fontsize=9.5, color="#64748b") # ── (b) 相图 ── ax = axes[1] style(ax, "(b) 存活 / 死亡的相图", "λ / σ²(该维的信噪比)", "β") snr = np.array(BK.LAMS) / s2 grid = np.logspace(-1, 1.6, 60) ax.loglog(grid, grid, "-", color="#94a3b8", lw=2.2, label="理论分界:β = λ/σ²") ax.fill_between(grid, grid, 1e3, color=C_BAD, alpha=0.06) ax.fill_between(grid, 1e-3, grid, color=C_OK, alpha=0.06) for lam in BK.LAMS: for beta in BK.BETAS: a, b, s, mse, kl, au = table[(lam, beta)] ok = BK.active_unit(a, lam) ax.plot(lam / s2, beta, "o" if ok else "s", color=C_OK if ok else C_BAD, markersize=7, mfc=C_OK if ok else "white", mew=1.6) ax.text(0.05, 0.93, "上方:KL 太贵,维度被关掉", transform=ax.transAxes, fontsize=10, color=C_BAD, weight="bold") ax.text(0.05, 0.06, "下方:信息量盖过 KL 的价,维度存活", transform=ax.transAxes, fontsize=10, color=C_OK, weight="bold") ax.set_xlim(0.1, 40) ax.set_ylim(0.15, 20) ax.legend(fontsize=9.5, frameon=False, loc="center right") ax.grid(axis="y", color="#eef2f7", lw=1.0) fig.suptitle("这张图要看什么:β 不是在「调重建和 KL 的比例」," "是在给每个潜变量维度标一个价——信噪比不够的维度直接归零", fontsize=11.5, color="#64748b", y=1.02) fig.tight_layout() p = OUT / "beta_tradeoff.png" fig.savefig(p, bbox_inches="tight", facecolor="white") plt.close(fig) print(f" ✓ {p.name}") if __name__ == "__main__": print("生成配图 →", OUT) fig_elbo_gap() fig_grad_variance() fig_beta_tradeoff() print("完成:3 张")
2026年09月25日
2 阅读
0 评论
0 点赞
2026-09-25
AIGC 基本功|FlashAttention 为什么不需要存下注意力矩阵-FlashAttn
FlashAttention 为什么不需要存下注意力矩阵 所属方向:注意力与位置编码 | 难度:进阶 | 前置知识:自注意力机制的计算与显存账本(本篇是它的直接后继) 关键词:FlashAttention、online softmax、tiling、IO感知、显存优化 01. 为什么需要它 上一篇《自注意力机制的计算与显存账本》结尾留了一个没解的结:N=4096 的视频 DiT,按保守教学分配模型计,单层注意力为 2.39 GiB 激活,无重算训练时按同一假设累加 32 层就是 76.5 GiB,尚未包含权重、梯度、优化器与 FFN。推理的临时激活则不能直接乘层数。当时摆出了两条出路:稀疏化(算得少,但要看清赔的是什么质量)和 FlashAttention。这一篇把后者讲透。 先把一个流传很广的说法钉死:FlashAttention 不是近似注意力。它算出来的就是标准的 softmax 注意力,在实数算术下等价,浮点下允许舍入差异(第 04 节有实测:float64 下最大误差 6.1e-16,纯浮点舍入)。它快的理由也不神秘——同目录 io_ledger.py 算过一笔账,单头 d=64、N=4096 时: 朴素实现要在 HBM(显存)上搬 130 MiB:分数矩阵 S 写一次读一次,softmax 权重 P 写一次读一次,4N² 次元素搬运; FlashAttention 只搬 17.8 MiB,少了 7.3 倍。 两种实现的主导矩阵乘 FLOPs 相同,但重算和归一化开销不同。在这组 A100 参数的 roofline 模型中,朴素实现算术强度为 31.5 FLOP/byte,低于约 200 的平衡点;这支持优先优化 IO 的方向,不证明所有序列、硬件或近似注意力方法都必然更慢。本文 130 MiB、17.8 MiB 是脚本估算值,不是 GPU 访存计数器实测。 另一个结论更值钱:显式保存概率矩阵 P 的朴素实现仍需二次存储,S 本身通常不必一并保留,显存永远是二次的;FlashAttention 把这两个矩阵整个从显存里删掉了,训练激活从 O(N²) 降到 O(N)——在统一采用六份线性张量、忽略小的 LSE 与工作区的教学预算中,上面那个 76.5 GiB 降为 4.5 GiB。N=32768 的长视频任务,本文教学模型的朴素实现要约 4.54 TiB 激活,物理上不存在能装下的卡;FlashAttention 只要 36 GiB。整个长序列时代(32K、128K 上下文)就是踩在这个技巧上站起来的。 02. 最小可用理解 三句话讲完核心思想: softmax attention 每行需要两个标量和一个 d 维向量:这行的最大值 m、指数和 l、加权和 O。把分数按块流进来,每块用一条递推式把这三个量修正一次,全流完之后 O÷l 就是精确的 softmax 注意力输出——中间任何一步都不需要把整行摆在内存里。这叫 online softmax。 显式物化 N×N 经常造成较低的算术强度;受限类型取决于维度、硬件与实现。朴素实现把 N×N 的 S、P 写回显存再读回来;FlashAttention 用 tiling(分块)把它们关在片上 SRAM 里算完就扔,HBM 读写量从 Θ(N²) 降到 Θ(N²d²/M),M 是片上 SRAM 的大小。 反向传播要用的 S、P 全部重算。除 Q/K/V(或它们的重算来源)外,前向存输出 O 和每行的 logsumexp(都是 O(N)),反向时拿 Q、K 重新过一遍分块流程,把需要的局部 P 重新算出来——用重复计算换显存,这是整个方法里最「敢」的一步。 这张图要看什么:左边朴素实现的 S、P 两个 N×N 是该前向示意中的主要中间量,它们必须写回 HBM 再读回来;右边是同一个 N×N 被切成 B_r×B_c 的小块,K/V 块进 SRAM 常驻、Q 块逐行流过,跨块只有 O_i、l_i、m_i 三个 O(N) 的量一直活着。 03. 数学推导 3.1 出发点:safe softmax 为什么需要先看全一行 设一行分数为 $s_1, \dots, s_N$,softmax 的定义是 $$p_i = \frac{\exp(s_i)}{\sum_{j=1}^{N} \exp(s_j)}$$ 分子分母都是 exp 的和。直接算会溢出:s 只要有 89 左右,$\exp(s)$ 在 fp32 就到 inf 了。工程上全部改用 safe softmax——先求这行的最大值 m,再算平移后的指数: $$p_i = \frac{\exp(s_i - m)}{\sum_{j=1}^{N} \exp(s_j - m)}, \qquad m = \max_{1 \le j \le N} s_j$$ 原始分子分母同乘 $\exp(-m)$,结果不变,但指数的输入全部落在 $(-\infty, 0]$,永不溢出。问题就出在这个 m 上:m 是对整行取的 max。你必须先把 N 个分数全部看过一遍才知道 m 是多少,然后才能开始算 exp——这解释了常见实现先保存分数再做 softmax 的流程;算法并不强制保存整个 N×N,也可重算分数或逐行处理,只是 IO 和效率不同。 而注意力输出对这一行还要再多两个量:分母 $l = \sum_j \exp(s_j - m)$,以及加权和 $O = \sum_j \exp(s_j - m)\, v_j$($v_j$ 是第 j 个 token 的 Value 向量)。最终输出就是 $O / l$。 所以真正要回答的问题是:如果分数是一块一块到来的(事先不知道后面块里有什么),这三个量还能算吗? 3.2 online softmax 递推式(全文核心) 能。做法是把「以 m 为参考系」改成「以当前的 m 为参考系,m 变了就整体换算」。 设已经流过了前 t 块,维护三个量:参考系最大值 $m^{(t)}$、分母 $l^{(t)}$、未归一化加权和 $O^{(t)}$,它们满足不变式 $$O^{(t)} = \sum_{j \le t} \exp(s_j - m^{(t)})\, v_j, \qquad l^{(t)} = \sum_{j \le t} \exp(s_j - m^{(t)})$$ (这里 $j \le t$ 是「属于前 t 块的所有下标」的缩写。)现在第 $t+1$ 块到了,块内最大值是 $m_{\text{blk}}$。新的全局最大值是 $$m^{(t+1)} = \max(m^{(t)},\ m_{\text{blk}})$$ 关键一步来了:旧累积量是按 $m^{(t)}$ 为参考系记的,而新的不变式要求参考系换成 $m^{(t+1)}$。把不变式里的 $\exp(s_j - m^{(t)})$ 拆成 $\exp(s_j - m^{(t+1)}) \cdot \exp(m^{(t+1)} - m^{(t)})$,旧量的换算系数就是 $\exp(m^{(t)} - m^{(t+1)})$: $$l^{(t+1)} = \exp(m^{(t)} - m^{(t+1)})\, l^{(t)} + \sum_{j \in \text{blk}} \exp(s_j - m^{(t+1)})$$ $$O^{(t+1)} = \exp(m^{(t)} - m^{(t+1)})\, O^{(t)} + \sum_{j \in \text{blk}} \exp(s_j - m^{(t+1)})\, v_j$$ 每一步只是把指数拆成两项相乘再重新合并,等价性是代入即可验证的恒等式;所有块流完后 $O^{(T)}/l^{(T)}$ 与朴素 softmax 在实数算术下相同(差在浮点舍入,第 04 节实测 1e-16 量级)。这个「换参考系」的系数在论文和代码里叫 rescale,跨块最大值、求和与 rescale 都会带来额外的标量操作;它们不改变主导矩阵乘次数。 严谨一点可以正向验证不变式:假设第 t 步的不变式成立,那么 $$O^{(t+1)} = e^{m^{(t)} - m^{(t+1)}} \sum_{j \le t} e^{s_j - m^{(t)}} v_j + \sum_{j \in \text{blk}} e^{s_j - m^{(t+1)}} v_j = \sum_{j \le t+1} e^{s_j - m^{(t+1)}} v_j$$ (第一个等号就是递推式,第二个等号把 $e^{m^{(t)} - m^{(t+1)}}$ 乘进求和号里、指数相加后正好变回 $e^{s_j - m^{(t+1)}}$。)旧块和新块在同一个参考系下合并,不变式保持。$l$ 的证明一字不差,把 $v_j$ 去掉就行。归纳基础是初始状态 $m^{(0)} = -\infty$、$l^{(0)} = 0$、$O^{(0)} = 0$:第一块到来时换算系数按 0 处理(对应代码里 np.where(np.isneginf(m_old), 0.0, ...) 那一行),三个量直接等于第一块的局部值。 三个细节值得停一下: 这个递推对任意分块都成立,块大小可以是 1(一个 key 一个 key 地流),也可以是 128 行。块大小只影响效率,不影响结果。 $\exp(m^{(t)} - m^{(t+1)}) \le 1$ 恒成立(因为 $m$ 单调不减),rescale 永远是在把旧量缩小,数值上很安全。 $m^{(t)} + \log l^{(t)}$ 就是 logsumexp(LSE)。它是反向传播唯一需要额外记录的东西——记住这个,第 3.3 节要用。 顺带说一句因果掩码:掩码就是把被遮位置的 $s$ 设成 $-\infty$,exp 之后是 0,对 m、l、O 都没有贡献;更进一步,如果一整块都被遮住(Q 块整体在 K 块之前),这块连算都不用算,直接跳过。第 04 节实测这个「整块跳过」在 N=4096、块 64 时省掉 49.2% 的块。 3.3 反向传播:重算换显存 这里的 O 指最终归一化输出。令 $G=\partial L/\partial O$、$A=GV^\top=\partial L/\partial P$。softmax 沿每行归一化,其正确反向公式是: $$\frac{\partial L}{\partial S}=P\odot\left(A-\operatorname{rowsum}(A\odot P)\right)$$ 行和为 $N\times1$,沿 key 轴广播;它是上游梯度在概率权重下的行平均,不能写成 $1-P^\top\mathbf 1$。随后 $dQ=dS\,K/\sqrt d$、$dK=dS^\top Q/\sqrt d$、$dV=P^\top G$。朴素实现可保存 P 而无需同时保存 S;FlashAttention 则重算局部 P。 FlashAttention 的做法是:除 Q/K/V 外,前向额外存 O 和 LSE($N \times d$ 加 $N$ 个数,O(N));反向时把分块流程原样再走一遍,在每一块里用 $\exp(S_{ij} - \text{LSE})$ 把局部的那一小块 P 重新算出来,立刻用于梯度,算完就扔。整个反向里 P 从头到尾没有以 N×N 的形态存在过。 用重复计算换显存——这笔交易的换算率是:每层每个头多算一遍 QK^T 和一次 exp(具体比例取决于前后向统计口径与实现),对本文保守教学账本,这对应去掉 S、P、浮点 dropout 乘子三项;真实朴素反向通常不用同时保存 S,具体减少几份取决于实现。N=4096、本文简化模型上,单层激活 2.39 GiB → 0.141 GiB,17 倍。 3.4 IO 账:搬运量到底差多少 设单头维度 d,片上存储预算为 M 个元素(不是字节);fp16 下 192 KiB 对应 M=98304。SRAM 里要同时放下 K 块、V 块(各 $B_c \times d$)和 Q 块、输出块(各 $B_r \times d$),论文取 $$B_c = \frac{M}{4d}, \qquad B_r = \min(B_c,\ d)$$ A100 每个 SM 有 192 KB SRAM,fp16 下 d=64 时 $B_c = 384$、$B_r = 64$——这是论文 IO 模型的粗略分块预算;真实 kernel 还受 score tile、累加器、寄存器、共享内存配额与 occupancy 约束,192 KiB 也不是每个 block 可独占的共享内存。 朴素实现的搬运量(单位:元素个数):QK^T 读 Q、K 各 Nd、写 S 一次 N²;softmax 读 S 写 P 各 N²;PV 读 P 一次 N²、读 V 写 O 各 Nd。合计 $4Nd + 4N^2$,主导项 4N²。 FlashAttention 的搬运量:K、V 各进 SRAM 一次($2Nd$);Q 块每换一个 K/V 块就要重读一遍,共 $T_c \cdot Nd$($T_c = \lceil N/B_c\rceil$ 是 K/V 块数);输出 O 同理要读出写回各一遍($2 T_c Nd$);l、m 两个 O(N) 的运行量共 $4T_c N$。合计约 $2Nd + 3T_cNd + 4T_cN$,代进去: $$\text{HBM 搬运量} \;\approx\; 3 \cdot \frac{N}{B_c} \cdot Nd \;=\; \Theta\!\left(\frac{N^2 d^2}{M}\right)$$ 两个量级一比:朴素是 $\Theta(N^2)$,Flash 是 $\Theta(N^2 d^2/M)$,比值 $d^2/M$ 在 d=64、fp16、M=98304 个元素(192 KiB) 时约等于 0.04——大 O 比例省略了常数,不能直接当作 25 倍的实际收益;脚本计入 Q/O 反复读写后的模型比值为 7.3 倍。这就是整篇论文的全部:不是新数学,是把「数据在哪」当成一等公民来优化。 按本文简化模型,朴素注意力的算术强度在 N 远大于 d 时约为 d/2 FLOP/byte(fp16),所以固定 d=64 时接近 32;它会随 d 改变,并非与头维度无关。A100 示例的约 200 FLOP/byte 是 dense FP16 Tensor Core 峰值与 HBM 带宽之比,真实 softmax 的非矩阵乘指令、缓存与调度仍会影响性能。 这张图要看什么:固定 d 与 SRAM 大小时,左图两者的大 N 主导项均为 N²,比值渐近趋于常数;线性项和取整会让有限 N 的比值变化;右图斜率不同(N² 对 N),N 越大两条线离得越远,朴素教学激活曲线越过示例 80 GiB 预算线的位置,FlashAttention 还有几百倍余量。 04. 代码实现 完整脚本在文末附录(flash_online_softmax.py、io_ledger.py、make_figures.py),只依赖 numpy,全部用 /usr/local/bin/python3 实跑过,下面的数字都是真实输出。 4.1 递推式长什么样:逐块打印 先看一个能盯着看的例子(N=8、D=4、K/V 块大小 B_c=2,Q 整行一起处理): def attention_flash(Q, K, V, B_r, B_c): N, D = Q.shape scale = 1.0 / np.sqrt(D) O = np.zeros((N, D)) # 未归一化的输出累加器 l = np.zeros(N) # 归一化分母(exp 之和) m = np.full(N, -np.inf) # 到目前为止见过的最大值 for j0 in range(0, N, B_c): Kj = K[j0:j0 + B_c] # [B_c, D] Vj = V[j0:j0 + B_c] # [B_c, D] for i0 in range(0, N, B_r): Qi = Q[i0:i0 + B_r] # [B_r, D] Sij = (Qi @ Kj.T) * scale # 局部分数 [B_r, B_c];Pij 也是同尺寸临时量 m_blk = Sij.max(axis=-1) # [B_r] m_old = m[i0:i0 + B_r] m_new = np.maximum(m_old, m_blk) Pij = np.exp(Sij - m_new[:, None]) # [B_r, B_c] l_blk = Pij.sum(axis=-1) corr = np.where(np.isneginf(m_old), 0.0, np.exp(m_old - m_new)) l[i0:i0 + B_r] = l[i0:i0 + B_r] * corr + l_blk O[i0:i0 + B_r] = O[i0:i0 + B_r] * corr[:, None] + Pij @ Vj m[i0:i0 + B_r] = m_new return O / l[:, None], (m + np.log(l)) out, lse = attention_flash(Q, K, V, B_r=64, B_c=64) print(out.shape, lse.shape) # (N, D) (N,) —— 输出和 LSE 都只有 O(N) 和第 3.2 节逐符号对上:m_new 是 $m^{(t+1)}$,corr 是换参考系的 $\exp(m^{(t)} - m^{(t+1)})$,l_blk 是块内指数和,Pij @ Vj 是块内加权和。corr 里的 np.where(np.isneginf(m_old), 0.0, ...) 处理的是第一块之前 $m = -\infty$ 的情况($-\infty - (-\infty)$ 会出 nan,直接规定换算系数为 0,旧累积量本来就是 0)。最后一行把 $m + \log l$ 作为 LSE 返回,留给反向。 实跑的逐块轨迹(第 0 号 query): j= 0 m=+0.2968 l=1.7748 corr=0.0000 j= 2 m=+0.6071 l=3.0317 corr=0.7333 j= 4 m=+0.6071 l=3.9732 corr=1.0000 j= 6 m=+0.6071 l=5.3007 corr=1.0000 看两点:j=2 时新块里出现了更大的分数,m 被抬高、旧累积量被打了个 0.733 的折扣;j=4 之后 m 没再变,corr 恒为 1,rescale 白做——此时乘子在数学上为 1;具体 GPU kernel 是否跳过这些操作要核对实现,不能仅由轨迹推断。最大的单个临时矩阵只有 16 个元素([8, 2]),而 N×N 是 64;这里不等于同时存活临时数组的总元素数。 4.2 等价性:和朴素实现在实数算术下等价,浮点下允许舍入差异 attention_naive 是上一篇的标准实现(S、P 都落地),两者对拍: N D causal 最大绝对误差 相对误差 256 64 False 6.106e-16 1.311e-15 256 64 True 8.882e-16 3.520e-16 1024 64 False 6.106e-16 2.629e-15 1024 64 True 6.661e-16 2.332e-16 2048 64 False 9.437e-16 3.724e-15 2048 64 True 8.327e-16 3.570e-16 4096 64 False 7.702e-16 5.679e-15 4096 64 True 7.772e-16 2.426e-16 误差全是 1e-16 量级——float64 的舍入级别。这就是「精确注意力」四个字的实测含义:不是「误差很小」,是算法本身和朴素 softmax 完全等价。 4.3 峰值内存:O(N²) 对 O(N),实测 用 tracemalloc 量函数内新分配的峰值内存(float64、D=64、块 64×64;Q/K/V 已提前分配,不计入本表): N naive 实测 naive 理论 2N²·8 flash 实测 比值 1024 16.5 MiB 16.0 MiB 1.1 MiB 14.4x 2048 65.0 MiB 64.0 MiB 2.2 MiB 30.1x 4096 258.0 MiB 256.0 MiB 4.2 MiB 61.5x 8192 1028.0 MiB 1024.0 MiB 8.3 MiB 123.6x N 从 4096 翻到 8192:naive 峰值 ×4.0,flash 峰值 ×2.0 naive 的实测和理论列($2N^2 \times 8$ 字节,S 和 P 两个 N×N)对得上,说明量的方法可信。最后那行是阶数的直接证据:N 翻倍,naive 峰值 ×4(二次),flash 峰值 ×2(线性)。 4.4 时间:在 CPU 上它反而更慢——这一点必须诚实 N naive flash(B=64) flash/naive 1024 7.2 ms 10.2 ms 1.4x 2048 26.3 ms 42.4 ms 1.6x 4096 98.7 ms 166.3 ms 1.7x CPU + numpy 上 flash 慢约 1.6 倍。三个原因:乘加次数一样还多了 rescale 和逐块 exp;一次大矩阵乘被拆成 (N/64)² 个 64×64 小矩阵乘,BLAS 跑不满小块,Python 循环开销也进来了;CPU 也有 SRAM 缓存,但这里的 Python/NumPy 分块没有实现专门的缓存与线程优化,不能把它当作 GPU 内核速度的预测。第三张图展示 IO 模型为何支持在 GPU 上尝试这一优化。 4.5 因果掩码:整块跳过,白捡一半 朴素实现加因果掩码,N×N 还是得整块算完再往被遮的位置上写 $-\infty$,一个 FLOP 都省不下来。FlashAttention 按「Q 块整体在 K 块之前就整块跳过」处理,N=4096、块 64 时实测: 块总数(非因果) : 4096 块总数(因果跳过): 2080 跳过的块占比 : 49.2% 保留的是含对角线的下三角,块数是 $T(T+1)/2$,占比 $(T+1)/2T$,T 大时趋近一半。训练 GPT 类因果模型、以及视频 DiT 里的时序因果注意力,这半是免费的。 05. 工业级实现对照 最小实现讲清了原理,但生产 kernel 和它有四处本质差异,每处都值得知道为什么: 第一,块大小不是从公式算的,是 autotune 出来的。 第 3.4 节的 $B_c = M/4d$ 是IO 分析采用的可行块预算;真实的 flash-attention kernel(flash_attn/flash_attn_interface.py 的 flash_attn_func,以 2025-09 的实现为准)里,块大小是按(头维度、数据类型、是否因果、显存架构)在若干组预编译配置里选的,还受 warp 数量、寄存器压力、shared memory bank conflict 的影响——公式只负责告诉你「必须小于某个数」,调优负责在约束内找最快的。 第二,减少非矩阵乘操作。 FlashAttention-2 使用未归一化输出累计,减少 rescale、除法等非矩阵乘工作。corr=1 时数学上无需改变旧值,但不能笼统声称所有 kernel 都按每行最大值是否变化来分支跳过。 第三,前向的结构是「外层 Q、内层 K/V」。 我们按论文 v1 的写法外层遍历 K/V 块;FlashAttention-2 把循环反过来(外层 Q 块),好处是输出 O 常驻寄存器不用反复读写、且不同 Q 块之间天然并行,能吃满更多 SM。论文 v1 的伪代码适合理解递推,v2 的循环结构才是现在 kernel 的样子。 第四,dropout 不存掩码,存随机数种子。 朴素实现要为反向留一个 B×H×N×N 的 dropout 掩码;kernel 里只存 Philox 计数器的 seed 和 offset(几十字节),反向时用同一个种子重新生成同样的掩码。这是「重算换显存」哲学最极致的一次应用——连随机数本身都可以重算。 另外两条工程事实:PyTorch 2.0 起 F.scaled_dot_product_attention 会自动按(头维度、掩码、数据类型、硬件)在 flash / memory-efficient / math 三个后端里挑,你不写一行 CUDA 也在用它;论文报告的端到端收益是 BERT-large(seq 512)比 MLPerf 1.1 训练记录快 15%、GPT-2(seq 1K)快 3 倍、Long Range Arena(seq 1K-4K)快 2.4 倍——注意 seq 512 时只有 15%,因为那时注意力在整层里占比还小,收益随序列长度涨,这正是 IO 复杂度模型的预测。 排查问题时你会想知道「此刻到底在用哪个后端」。PyTorch 留了一个官方口子: from torch.nn.attention import sdpa_kernel, SDPBackend with sdpa_kernel(SDPBackend.FLASH_ATTENTION): out = F.scaled_dot_product_attention(q, k, v, is_causal=True) # 强制走 flash;如果这个头维度/掩码组合它不支持,这里会直接报错, # 而不是悄悄退回 math 后端——「悄悄降级」正是性能莫名掉一半时最该先查的事 这张图要看什么:横轴是 N,纵轴是「每搬 1 字节做多少次运算」,灰色虚线是机器平衡点(201 FLOP/byte)。本模型中朴素实现的强度渐近接近 32;FlashAttention 的估计从 N=2048 起超过平衡点。这是按矩阵乘峰值做的模型分类,真实瓶颈还受非矩阵乘指令、缓存、并行度和调度影响。 06. 代价与边界 FlashAttention 省下了 HBM 搬运和 N² 显存,赔进去的和没管住的也要说清楚。 代价:重计算和额外归一化操作。 反向要重算局部分数及概率,其中包含矩阵乘和逐元素运算,不能统称为固定 30% 的额外非矩阵乘开销。小序列的收益取决于 kernel、调度和硬件,没有统一的 seq<512 亏损阈值。 数值边界:实数算术等价不保证浮点逐位一致。 kernel 内部用 fp16/bf16 存储、fp32 累加,块内的归一化和朴素实现的一次性归一化在浮点上不同。对训练的影响需要结合 dtype、输入尺度与误差测试判断,但如果你在做数值敏感的分析(比如逐 token 概率对比),要知道它和参考实现差在舍入级别,不是 bug。 边界的核心一条:它没有改变复杂度,改变的常数。 显存从 $O(N^2)$ 降到 $O(N)$,但算力还是 $\Theta(N^2 d)$、HBM 搬运还是 $\Theta(N^2 d^2/M)$。N=32768 时 FlashAttention 的单头搬运是 1061 MiB——比朴素实现的 8.2 GiB 好得多,但随 N 继续平方增长这一点没变。上下文再往上涨(1M token),接力棒要交给稀疏注意力、线性注意力、状态空间模型这些真正改复杂度的方法。FlashAttention 的块结构恰恰是它们的底座:把注意力切成块之后,「整块跳过」才成为可能,第 4.5 节那个 49.2% 推广到任意稀疏模式就是块稀疏注意力。 不该用的场景:需要拿到完整注意力权重做分析或可视化的(P 从头到尾没存在过,想看它就得回到朴素实现);自定义的任意注意力偏置如果 kernel 不支持,绕过去的方法可能把优势吃掉;以及缺乏适配内核的环境;CPU 也有 SRAM 缓存,但本文 NumPy 循环没有实现专门的 CPU cache 优化(第 4.4 节的 CPU 实测就是例子)。 07. 经典论文脉络 这条线的演进关系一句话各说清: Milakov & Gimelshein, 2018(arXiv:1805.02867)Online normalizer calculation for softmax——首次提出 softmax 的 online 计算:流式更新 max 和指数和。当时的目标只是省一次对 logits 的遍历,还没人把它和注意力显存联系起来。 Rabe & Staats, 2021(arXiv:2112.05682)Self-attention Does Not Need O(n²) Memory——讨论单 query 的常数额外空间方案,并给出分块的低内存自注意力实现及速度/内存实验;不能把单 query 的额外空间结论写成整段输出总存储 O(1),也不能概括为没有实用价值。 Dao et al., 2022(arXiv:2205.14135)FlashAttention——本篇锚点。补上缺失的一环:把流式更新从「逐 token」改成「逐块」(tiling),配上 IO 复杂度分析和 GPU kernel,第一次让「精确 + 更快 + 更省显存」三者同时成立。 Dao, 2023(arXiv:2307.08691)FlashAttention-2——循环重排(外层 Q)、削减非矩阵乘 FLOPs、更好的并行度,把 v1 大约 25-40% 的峰值算力利用率推向 50-73%。 一条清晰的线:2018 年有技巧,2021 年有证明,2022 年才有产品——缺的从来不是数学,是「意识到瓶颈在 IO」这个视角。 08. 常见误解 以下几条都值得单独记住,前两条我当初也信过: 「FlashAttention 是近似注意力,所以有精度损失」。错。它是精确算法,和朴素 softmax 数学等价(第 4.2 节实测 1e-16)。真正近似的是 Linformer、Performer 那一族。两者经常被并列讨论,但一个在改算法,一个在改数据的搬运方式。 「它快是因为算了更少的 FLOPs」。反了,它的 FLOPs 略多于朴素实现(第 4.4 节 CPU 实测慢 1.6 倍)。主要收益来自减少 IO,也来自更好的并行划分与调度——这正是它给所有做系统优化的人的启示:先问数据在哪,再问算了多少。 「显存优化只对前向有用」。最大的收益在反向:除 Q/K/V 外额外存 O 和 LSE(固定 d 时为 O(N)),反向重算 P。本文三份二次张量是教学分配假设,真实反向不一定保存三份;可靠结论是分块重算避免显式保留完整 P。 「块开得越大越快」。块大小受 SRAM 硬约束:超过每个 block 的资源上限可能无法启动或编译;寄存器 spill 等情况也会引入额外访存,整个方法的根基(中间结果不落 HBM)就塌了。真实 kernel 的块大小是约束内的调优问题,不是越大越好。 「有了它就不用再关心 N²」。它优化的是常数,不是阶数——算力和搬运仍随 N 平方涨。能继续走的长上下文路线是稀疏化(改算力阶数)和状态空间/线性注意力(改注意力本身),FlashAttention 的块结构是它们的载体,不是替代品。 09. 动手验证 跑文末附录的 flash_online_softmax.py(只依赖 numpy),预期输出与正文一致: 第 2 节等价性表:非因果最大误差 6.1e-16 到 9.4e-16,因果 6.7e-16 到 8.9e-16——全是舍入级别; 第 3 节内存表:naive 峰值和理论值 $2N^2 \times 8$ 字节吻合,N 翻倍时 naive ×4.0、flash ×2.0。 再做一个一行代码的实验,直接看清递推式里 rescale 的分量:把 corr = np.where(np.isneginf(m_old), 0.0, np.exp(m_old - m_new)) 改成 corr = np.ones_like(m_old)(假装 m 永远不变,也就是退回「先见全家再算」之前的朴素流式假设),重跑等价性检验。我实测过:N=1024、D=64、随机高斯输入下,最大误差从 6.1e-16 恶化到 0.276,平均误差 0.0136——输出在量级上就是错的。这个对比说明:流式计算 softmax 时,「用新参考系换算旧累积量」这一步不是工程细节,是正确性本身。 最后改 io_ledger.py 开头的硬件常数(比如把 SRAM_PER_SM 调到 48 KB 模拟消费级卡),重跑看块大小和搬运量比值怎么变——你会看到 SRAM 越小,FlashAttention 相对朴素实现的搬运量优势越小,$d^2/M$ 里的 M 直接控制这一切。 10. 延伸阅读 按知识树的依赖关系,建议按这个顺序继续走: 前置:自注意力机制的计算与显存账本——本篇所有显存账本的出处(2.39 GiB、94% 那几笔账都在那篇里);旋转位置编码 RoPE 与 视频 DiT 里的 3D RoPE——注意力的另外两个必备零件,本篇刻意没有碰位置编码。 后继:视频生成里的稀疏注意力(规划中)——FlashAttention 的块结构是稀疏模式的执行底座,「整块跳过」从因果掩码推广到任意稀疏图;算子融合与 CUDA Graph(规划中)——把「少搬内存」推到极端就是融合,FlashAttention 是这个思路最成功的案例;KV Cache 与自回归视频生成——推理时的另一本显存账。 延伸到知识树之外:想读 kernel 源码,从 Dao-AILab/flash-attention 的 flash_attn/flash_attn_interface.py 进,先读 forward 再读 backward;想读原始推导,Milakov & Gimelshein(1805.02867)给出了独立 softmax 在线归一化的直接推导。 附录:完整代码 09 节用到的脚本全文如下(io_ledger.py、flash_online_softmax.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 io_ledger.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """IO 账本:朴素注意力和 FlashAttention 在 HBM 上到底搬了多少字节。 这篇讲的是「为什么快」。答案不在 FLOPs 上——两者的乘加次数几乎一样—— 而在于 GPU 有两级内存: HBM(显存) 带宽约 1.5 TB/s,容量 40~80 GB SRAM(片上) 带宽约 19 TB/s,但每个 SM 只有 192 KB 朴素实现把 N×N 的分数矩阵写回 HBM 再读出来,等于把数据在这条 13 倍带宽差的 通道上来回搬;FlashAttention 用 tiling 让这些中间结果根本不落 HBM。 下面的模型把每一笔读写都点清楚,参数是可调的,读者可以改硬件常数重算。 硬件数量级取自 FlashAttention 论文 Table 1(A100 40GB)。 只依赖 numpy。直接 `python io_ledger.py` 即可运行。 """ import numpy as np # ── 硬件常数(A100 40GB 量级,论文 Table 1)──────────────────── HBM_BW = 1.555e12 # HBM 带宽,字节/秒 SRAM_BW = 19.0e12 # 片上 SRAM 带宽,字节/秒 SRAM_PER_SM = 192 * 1024 # 每个 SM 的片上 SRAM,字节 FLOPS_PEAK = 312e12 # fp16 tensor core 峰值,FLOP/秒 GIB = 1024 ** 3 MIB = 1024 ** 2 # ──────────────────────────────────────────────────────────── # 块大小:SRAM 里能同时放下什么 # ──────────────────────────────────────────────────────────── def block_sizes(d, sram_bytes=SRAM_PER_SM, dtype_bytes=2): """返回 (B_r, B_c):Q 块行数与 K/V 块行数。 SRAM 里要同时放下 K_j、V_j 两个 [B_c, d] 和 Q_i、O_i 两个 [B_r, d]。 论文的取法是 B_c = M / (4d)、B_r = min(B_c, d),这里照抄: 先让 K_j+V_j 占掉一半 SRAM,Q 块则不超过 d 行(保证 softmax 按行算得下)。 """ elems = sram_bytes / dtype_bytes B_c = max(1, int(elems // (4 * d))) B_r = min(B_c, d) return B_r, B_c # ──────────────────────────────────────────────────────────── # HBM 读写量(单位:元素个数,乘 dtype_bytes 得字节) # ──────────────────────────────────────────────────────────── def hbm_elems_naive(N, d): """朴素实现:S 和 P 都要落地。 QK^T: 读 Q(Nd) + 读 K(Nd) + 写 S(N²) softmax: 读 S(N²) + 写 P(N²) PV: 读 P(N²) + 读 V(Nd) + 写 O(Nd) """ return 4 * N * d + 4 * N * N def hbm_elems_flash(N, d, B_r, B_c): """FlashAttention:外层遍历 K/V 块,内层遍历 Q 块。 K、V 各读一遍(每个 j 块进 SRAM 后,内层 i 循环里一直复用) Q 每个 j 都要重读一遍:T_c · N·d O 每个 (j,i) 都要读出来再写回去(累加器跨 j 迭代):2 · T_c · N·d l、m 两个 O(N) 的运行量同理:4 · T_c · N """ T_c = int(np.ceil(N / B_c)) return 2 * N * d + 3 * T_c * N * d + 4 * T_c * N def flops_attention(N, d): """两个 N×N 矩阵乘,一次 [M,K]×[K,N] 算 2MKN 个浮点运算。""" return 4 * N * N * d def activation_bytes(N, d, H, B=1, dtype_bytes=2): """单层注意力的训练激活(反向要用的中间张量)。 保守教学模型:6 个 [B,N,d_model] 线性项 + S/P/浮点乘子 3 个 [B,H,N,N]。 此函数 d 是总宽度 d_model,不是其他 IO 函数中的头宽度。 融合侧保留同样六份线性预算,忽略小的 LSE 与工作区;不是框架峰值测量。 """ lin = 6 * B * N * d * dtype_bytes quad = 3 * B * H * N * N * dtype_bytes return lin + quad, lin def roofline(bytes_moved, flops): """算术强度(FLOP/byte)与两个上界时间。返回 (强度, 内存时间, 算力时间, 瓶颈)。""" intensity = flops / bytes_moved t_mem = bytes_moved / HBM_BW t_comp = flops / FLOPS_PEAK bound = "内存受限" if t_mem > t_comp else "算力受限" return intensity, t_mem, t_comp, bound # ──────────────────────────────────────────────────────────── # 报表 # ──────────────────────────────────────────────────────────── def report_io(): d, b = 64, 2 # 单头维度 64,fp16 B_r, B_c = block_sizes(d) print("=" * 74) print("1. SRAM 块大小与 HBM 读写量(单头,d=64,fp16)") print("=" * 74) print(f" SRAM {SRAM_PER_SM / 1024:.0f} KB / SM,fp16 下能放 " f"{SRAM_PER_SM / b:.0f} 个元素") print(f" → B_c = M/(4d) = {B_c},B_r = min(B_c, d) = {B_r}\n") print(f"{'N':>7} {'naive HBM':>12} {'flash HBM':>12} {'比值':>8} " f"{'naive 强度':>11} {'flash 强度':>11}") rows = [] for N in (1024, 2048, 4096, 8192, 16384, 32768): nb = hbm_elems_naive(N, d) * b fb = hbm_elems_flash(N, d, B_r, B_c) * b fl = flops_attention(N, d) i_n, _, _, bound_n = roofline(nb, fl) i_f, _, _, bound_f = roofline(fb, fl) rows.append((N, nb, fb, i_n, i_f, bound_n, bound_f)) print(f"{N:>7} {nb / MIB:>10.1f} MiB {fb / MIB:>10.1f} MiB " f"{nb / fb:>7.1f}x {i_n:>9.1f} {i_f:>9.1f} ") print(f"\n 机器平衡点(峰值算力/带宽)= {FLOPS_PEAK / HBM_BW:.0f} FLOP/byte") print(" 强度低于它 → 内存受限,加算力没用;高于它 → 才开始吃算力。\n") print(f"{'N':>7} {'naive 瓶颈':>12} {'flash 瓶颈':>12}") for N, nb, fb, i_n, i_f, bn, bf in rows: print(f"{N:>7} {bn:>12} {bf:>12}") print("\n 注意:flash 的强度在 N 大时越过平衡点,模型说它变成算力受限了。") print(" 但真实 kernel 达不到峰值——softmax 的 exp 走的是特殊函数单元,") print(" 不走 tensor core,这个「非矩阵乘开销」正是 FlashAttention-2 之后") print(" 继续优化的地方。模型给出的是上界,不是承诺。") def report_activation(): d, H, B, b = 3072, 24, 1, 2 print("\n" + "=" * 74) print("2. 教学激活存储模型(简化 Transformer:d=3072, H=24, B=1, 32 层, bf16)") print("=" * 74) print(f"{'N':>7} {'朴素/层':>12} {'Flash/层':>12} {'比值':>8} " f"{'朴素 32 层':>12} {'Flash 32 层':>13}") for N in (1024, 4096, 8192, 32768): naive, flash = activation_bytes(N, d, H, B, b) print(f"{N:>7} {naive / GIB:>10.2f} GiB {flash / GIB:>10.3f} GiB " f"{naive / flash:>7.0f}x {32 * naive / GIB:>10.1f} GiB " f"{32 * flash / GIB:>11.2f} GiB") n4096, f4096 = activation_bytes(4096, d, H, B, b) print(f"\n N=4096 时,朴素实现单层 {n4096 / GIB:.2f} GiB,其中二次项占 " f"{100 * (1 - f4096 / n4096):.1f}%——") print(" 这正是 attention_basics 那篇里「三个 N×N 吃掉 94%」的那一格。") print(" FlashAttention 删掉的就是这一格,剩下的是 O(N) 的线性项。") def report_speed_limit(): d, b = 64, 2 B_r, B_c = block_sizes(d) print("\n" + "=" * 74) print("3. 理论上界:如果只受 HBM 带宽限制,两者各要多久(单头,d=64)") print("=" * 74) print(f"{'N':>7} {'FLOPs':>10} {'naive 内存时间':>16} {'flash 内存时间':>16} " f"{'纯带宽模型比':>9}") for N in (1024, 4096, 16384): nb = hbm_elems_naive(N, d) * b fb = hbm_elems_flash(N, d, B_r, B_c) * b fl = flops_attention(N, d) print(f"{N:>7} {fl / 1e9:>8.2f} G {nb / HBM_BW * 1e3:>14.3f} ms " f"{fb / HBM_BW * 1e3:>14.3f} ms {nb / fb:>8.1f}x") print("\n 此列只比较理想 HBM 时间,并非完整模型训练或实际 kernel 加速上界。") print(" 论文的端到端训练加速与此处单头 IO 模型口径不同,不能直接相比。") if __name__ == "__main__": report_io() report_activation() report_speed_limit() flash_online_softmax.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """online softmax 的正确性与显存代价实测(FlashAttention 的核心那一步)。 对比两种实现: naive —— 先算出完整的 N×N 分数矩阵 S,再逐行 softmax 得到 P,最后算 O = P V flash —— 按块流过 K/V,只维护三个 O(N) 的运行量:输出累加器 O、分母 l、 运行最大值 m(第 03 节递推式的直接实现) 两者在数学上完全等价,唯一差别是中间有没有 N×N 的矩阵落到内存里。 这份脚本回答三个问题: 1. 递推式写出来的结果,和朴素 softmax 逐位一致吗?(1.1 节) 2. 峰值内存真的差一个 N 吗?(用 tracemalloc 量,不是估的) 3. 那算力呢?——在 CPU + numpy 上 flash 是**更慢**的,这一点必须诚实讲清楚 只依赖 numpy。直接 `python flash_online_softmax.py` 即可运行。 """ import time import tracemalloc import numpy as np NEG = -np.inf # ──────────────────────────────────────────────────────────── # 两个被测实现 # ──────────────────────────────────────────────────────────── def attention_naive(Q, K, V, causal=False): """标准实现:S 和 P 都是完整的 N×N 常驻张量。""" D = Q.shape[-1] S = (Q @ K.T) / np.sqrt(D) # [N, N] ← 第一块 N×N if causal: S = np.where(np.triu(np.ones((Q.shape[0], K.shape[0])), 1) > 0, NEG, S) S -= S.max(axis=-1, keepdims=True) # safe softmax,不改变结果 P = np.exp(S) # [N, N] ← 第二块 N×N P /= P.sum(axis=-1, keepdims=True) return P @ V def attention_flash(Q, K, V, B_r, B_c, causal=False, trace=False): """按块流过 + online softmax。全程不出现 N×N 的张量。 B_r / B_c 分别是 Q 块和 K/V 块的行数,对应 SRAM 里各放得下多少行。 """ N, D = Q.shape scale = 1.0 / np.sqrt(D) O = np.zeros((N, D)) # 未归一化的输出累加器 l = np.zeros(N) # 归一化分母(exp 之和) m = np.full(N, NEG) # 到目前为止见过的最大值 peak_tmp = 0 # 记录出现过的最大临时矩阵(元素个数) for j0 in range(0, N, B_c): Kj = K[j0:j0 + B_c] # [B_c, D] Vj = V[j0:j0 + B_c] # [B_c, D] for i0 in range(0, N, B_r): # 因果掩码下,若整个 Q 块都在 K 块之前(所有 query 下标 < 所有 key # 下标),这一块全被遮掉,连算都不用算 —— 朴素实现做不到这一点。 # 条件:块内最大 query 下标 i0+B_r-1 < j0 if causal and i0 + B_r <= j0: continue Qi = Q[i0:i0 + B_r] # [B_r, D] Sij = (Qi @ Kj.T) * scale # [B_r, B_c] 局部分数;Pij 也是临时块 if causal: q_idx = i0 + np.arange(Sij.shape[0])[:, None] k_idx = j0 + np.arange(Sij.shape[1])[None, :] Sij = np.where(k_idx > q_idx, NEG, Sij) peak_tmp = max(peak_tmp, Sij.size) m_blk = Sij.max(axis=-1) # [B_r] m_old = m[i0:i0 + B_r] m_new = np.maximum(m_old, m_blk) # [B_r] Pij = np.exp(Sij - m_new[:, None]) # [B_r, B_c] l_blk = Pij.sum(axis=-1) # [B_r] # 把之前累积的量从旧的参考最大值搬到新的(关键的 rescale 一步) corr = np.where(np.isneginf(m_old), 0.0, np.exp(m_old - m_new)) l[i0:i0 + B_r] = l_old = l[i0:i0 + B_r] * corr + l_blk O[i0:i0 + B_r] = O[i0:i0 + B_r] * corr[:, None] + Pij @ Vj m[i0:i0 + B_r] = m_new if trace: print(f" j={j0:>2} i={i0:>2} m={m_new[0]:+.4f} " f"l={l_old[0]:.4f} corr={corr[0]:.4f}") return O / l[:, None], (m + np.log(l)), peak_tmp def peak_bytes(fn, *args, **kwargs): """跑一次 fn,返回 (结果, 峰值字节数)。用 tracemalloc 实测。""" tracemalloc.start() tracemalloc.reset_peak() out = fn(*args, **kwargs) _, peak = tracemalloc.get_traced_memory() tracemalloc.stop() return out, peak def timed(fn, *args, repeat=3, **kwargs): best = float("inf") out = None for _ in range(repeat): t0 = time.perf_counter() out = fn(*args, **kwargs) best = min(best, time.perf_counter() - t0) return out, best # ──────────────────────────────────────────────────────────── # 1. 递推式长什么样:一个能逐块打印的最小例子 # ──────────────────────────────────────────────────────────── def demo_recurrence(): rng = np.random.default_rng(0) N, D, B_c = 8, 4, 2 Q = rng.normal(size=(N, D)) K = rng.normal(size=(N, D)) V = rng.normal(size=(N, D)) print("=" * 68) print("1. online softmax 的递推过程(N=8, D=4, K/V 块大小 B_c=2)") print("=" * 68) print(" Q 块固定为整行(B_r=N),K/V 分成 4 块依次流过;") print(" 每行打印第 0 号 query 的 m / l / corr,看它们怎么被逐次修正:\n") _, _, peak = attention_flash(Q, K, V, B_r=N, B_c=B_c, trace=True) print(f"\n 最大的单个临时矩阵只有 {peak} 个元素(并非临时内存总和) = [{N}, {B_c}],而 N×N = {N * N}") # ──────────────────────────────────────────────────────────── # 2. 等价性:和朴素 softmax 逐位对得上吗 # ──────────────────────────────────────────────────────────── def demo_exactness(): print("\n" + "=" * 68) print("2. 数值等价性(float64,非因果 / 因果两种掩码)") print("=" * 68) print(f"{'N':>6} {'D':>4} {'causal':>7} {'最大绝对误差':>14} {'相对误差':>12}") for N in (256, 1024, 2048, 4096): D = 64 rng = np.random.default_rng(N) Q = rng.normal(size=(N, D)) K = rng.normal(size=(N, D)) V = rng.normal(size=(N, D)) for causal in (False, True): ref = attention_naive(Q, K, V, causal) got, _, _ = attention_flash(Q, K, V, B_r=64, B_c=64, causal=causal) diff = np.abs(ref - got).max() rel = diff / np.abs(ref).max() print(f"{N:>6} {D:>4} {str(causal):>7} {diff:>14.3e} {rel:>12.3e}") print("\n 误差量级是浮点舍入(1e-15),不是近似——FlashAttention 是精确算法。") # ──────────────────────────────────────────────────────────── # 3. 峰值内存:是不是真的差一个 N # ──────────────────────────────────────────────────────────── def demo_memory(): print("\n" + "=" * 68) print("3. 峰值内存实测(tracemalloc,float64,D=64,B_r=B_c=64)") print("=" * 68) print(f"{'N':>6} {'naive 实测':>12} {'naive 理论 2N²·8':>18} " f"{'flash 实测':>12} {'比值':>8}") peaks = {} for N in (1024, 2048, 4096, 8192): D = 64 rng = np.random.default_rng(N) Q = rng.normal(size=(N, D)) K = rng.normal(size=(N, D)) V = rng.normal(size=(N, D)) _, p_naive = peak_bytes(attention_naive, Q, K, V) _, p_flash = peak_bytes(attention_flash, Q, K, V, 64, 64) peaks[N] = (p_naive, p_flash) print(f"{N:>6} {p_naive / 2**20:>10.1f} MiB " f"{2 * N * N * 8 / 2**20:>16.1f} MiB " f"{p_flash / 2**20:>10.1f} MiB {p_naive / p_flash:>7.1f}x") g_n = peaks[8192][0] / peaks[4096][0] g_f = peaks[8192][1] / peaks[4096][1] print(f"\n N 从 4096 翻到 8192:naive 峰值 ×{g_n:.1f},flash 峰值 ×{g_f:.1f}") print(" 这就是 O(N²) 和 O(N) 的区别:前者翻两倍(4×),后者跟着翻倍(2×)。") print(" 朴素实现要同时留住 S 和 P 两个 N×N(理论列就是 2N²·8 字节,和实测对得上);") print(" flash 只留 B_r×B_c 的块,剩下的是 O(N·D) 的输出累加器。") # ──────────────────────────────────────────────────────────── # 4. 时间:CPU + numpy 上 flash 反而更慢,这才是重点 # ──────────────────────────────────────────────────────────── def demo_time(): print("\n" + "=" * 68) print("4. 墙钟时间(CPU + numpy,BLAS 多线程)——反直觉的一项是这个") print("=" * 68) print(f"{'N':>6} {'naive':>10} {'flash(B=64)':>13} {'flash/naive':>12}") for N in (1024, 2048, 4096): D = 64 rng = np.random.default_rng(N) Q = rng.normal(size=(N, D)) K = rng.normal(size=(N, D)) V = rng.normal(size=(N, D)) _, t_naive = timed(attention_naive, Q, K, V) _, t_flash = timed(attention_flash, Q, K, V, 64, 64) print(f"{N:>6} {t_naive * 1e3:>9.1f} ms {t_flash * 1e3:>12.1f} ms " f"{t_flash / t_naive:>11.1f}x") print("\n 以上比值是本机本次计时,不能据此预测 GPU;可能影响速度的因素包括:") print(" 1. 乘加次数一模一样,还额外多了每块的 rescale 和逐元素 exp;") print(" 2. 一次大矩阵乘被拆成 (N/B)² 次 64×64 的小矩阵乘,") print(" BLAS 在小块上根本跑不满,Python 循环开销也进来了;") print(" 3. 最关键的:它省的是**内存搬运**,不是 FLOPs,") print(" CPU 同样有 SRAM 缓存,但本示例没有专门优化缓存和线程,") print(" 因此分块可能省容量,却被 Python 循环与小 GEMM 开销抵消。") print(" 第 05 节讲 GPU 上为什么结论会反过来。") # ──────────────────────────────────────────────────────────── # 5. 因果掩码:flash 能顺手省掉一半算力,朴素实现不能 # ──────────────────────────────────────────────────────────── def demo_causal_blocks(): print("\n" + "=" * 68) print("5. 因果掩码下实际算了多少块(N=4096, B_r=B_c=64)") print("=" * 68) N, B = 4096, 64 T_r = T_c = N // B total = 0 for j in range(T_c): for i in range(T_r): if i * B + B <= j * B: # 整个 Q 块都在 K 块之前 → 全遮,跳过 continue total += 1 print(f" 块总数(非因果) : {T_r * T_c}") print(f" 块总数(因果跳过): {total}") print(f" 跳过的块占比 : {100.0 * (1 - total / (T_r * T_c)):.1f}%") print("\n 朴素实现就算加了因果掩码,N×N 的矩阵照样得整块算完再遮,") print(" 省不了一点算力;flash 是整块跳过,这是免费的一半(严格说是") print(" (T²+T)/2 / T² ≈ 一半多一点的对角块保留)。") if __name__ == "__main__": np.set_printoptions(precision=4, suppress=True) demo_recurrence() demo_exactness() demo_memory() demo_time() demo_causal_blocks() make_figures.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """生成「FlashAttention」的三张解释图。 数值全部来自同目录下的 io_ledger.py(HBM 读写量、激活显存、算术强度), 改了那个脚本的话这里要跟着重跑,避免图与正文数字不一致。 三张图分别回答: 1. tiling 到底怎么切、什么留在 SRAM、什么留在 HBM 2. 两本账随 N 怎么长(HBM 读写量、教学激活存储模型) 3. 为什么说是「内存受限」变的「不那么受限」(算术强度对机器平衡点) 只依赖 numpy + matplotlib,直接 `python make_figures.py` 即可运行。 """ from pathlib import Path import matplotlib.pyplot as plt import numpy as np from matplotlib.patches import FancyArrowPatch, FancyBboxPatch, Rectangle from io_ledger import ( FLOPS_PEAK, HBM_BW, activation_bytes, block_sizes, flops_attention, hbm_elems_flash, hbm_elems_naive, ) ROOT = Path(__file__).resolve().parents[1] OUT = ROOT / "figures" OUT.mkdir(exist_ok=True) plt.rcParams.update({ "font.sans-serif": ["PingFang SC", "Arial Unicode MS", "DejaVu Sans"], "axes.unicode_minus": False, "figure.dpi": 160, }) C_NAIVE = "#e05263" # 朴素实现:红 C_FLASH = "#2f9e6f" # FlashAttention:绿 C_SRAM = "#f0a03c" # SRAM 高亮:橙 C_GREY = "#94a3b8" C_SKIP = "#e4e9f0" INK = "#182238" MIB = 1024 ** 2 GIB = 1024 ** 3 def _style(ax, title, xlabel=None, ylabel=None): ax.set_title(title, fontsize=13.5, weight="bold", color=INK, pad=10) if xlabel: ax.set_xlabel(xlabel, fontsize=11, color="#475569") if ylabel: ax.set_ylabel(ylabel, fontsize=11, color="#475569") ax.tick_params(colors="#475569", labelsize=10) for s in ("top", "right"): ax.spines[s].set_visible(False) for s in ("left", "bottom"): ax.spines[s].set_color("#cbd5e1") ax.set_axisbelow(True) def _box(ax, x, y, w, h, text, fc, fs=10.5, tc="white"): ax.add_patch(FancyBboxPatch( (x, y), w, h, boxstyle="round,pad=0.02,rounding_size=0.12", linewidth=0, facecolor=fc, edgecolor="none")) ax.text(x + w / 2, y + h / 2, text, ha="center", va="center", fontsize=fs, color=tc, weight="bold", linespacing=1.5) def _arrow(ax, x1, y1, x2, y2, label=None, color="#475569"): ax.add_patch(FancyArrowPatch( (x1, y1), (x2, y2), arrowstyle="-|>", mutation_scale=13, linewidth=1.4, color=color, shrinkA=2, shrinkB=2)) if label: ax.text((x1 + x2) / 2, max(y1, y2) + 0.22, label, ha="center", fontsize=9, color=color) def _caption(fig, text): fig.text(0.5, 0.012, text, ha="center", fontsize=9.5, color="#94a3b8") # ──────────────────────────────────────────────────────────── # 图 1:tiling 怎么切 # ──────────────────────────────────────────────────────────── def fig_tiling(): fig = plt.figure(figsize=(13.6, 6.0), facecolor="#fbfcfe") # ── 左:朴素实现 ── ax = fig.add_axes([0.03, 0.09, 0.46, 0.80]) ax.set_xlim(0, 10); ax.set_ylim(0, 7.4); ax.axis("off") ax.text(0, 7.05, "本例朴素前向:物化 S 与 P", fontsize=13.5, weight="bold", color=C_NAIVE) _box(ax, 0.1, 5.1, 1.5, 1.1, "Q, K, V\n[N, d]", "#64748b", fs=10) _box(ax, 2.4, 4.8, 2.0, 1.7, "S = QK^T / √d\n[N, N]", C_NAIVE, fs=11) _box(ax, 5.3, 4.8, 2.0, 1.7, "P = softmax(S)\n[N, N]", C_NAIVE, fs=11) _box(ax, 8.2, 5.1, 1.6, 1.1, "O = PV\n[N, d]", "#64748b", fs=10) _arrow(ax, 1.65, 5.65, 2.35, 5.65) _arrow(ax, 4.45, 5.65, 5.25, 5.65, "读 S 写 P") _arrow(ax, 7.35, 5.65, 8.15, 5.65, "读 P") ax.text(4.4, 4.25, "N=4096, d=64, fp16:S 和 P 各 32 MiB", ha="center", fontsize=9.5, color="#475569") ax.add_patch(Rectangle((0.1, 2.2), 9.7, 1.35, facecolor="#eef2f7", edgecolor="#cbd5e1", linewidth=1.2)) ax.text(4.95, 3.2, "HBM 带宽 1.5 TB/s", ha="center", fontsize=10.5, color="#475569", weight="bold") ax.text(4.95, 2.5, "S 写一次读一次,P 写一次读一次 → 4N² 次元素搬运", ha="center", fontsize=9.5, color="#64748b") _arrow(ax, 3.4, 4.75, 3.4, 3.6, color=C_NAIVE) _arrow(ax, 6.3, 4.75, 6.3, 3.6, color=C_NAIVE) ax.text(0.1, 1.35, "代价:数据在这条通道上来回两趟,", fontsize=10.5, color=C_NAIVE, weight="bold") ax.text(0.1, 0.75, "而 HBM 比 SRAM 慢 12 倍", fontsize=10.5, color=C_NAIVE, weight="bold") # ── 右:FlashAttention ── ax2 = fig.add_axes([0.53, 0.09, 0.44, 0.80]) ax2.set_xlim(0, 10); ax2.set_ylim(0, 7.4); ax2.axis("off") ax2.text(0, 7.05, "FlashAttention:按块流过,只留 O(N)", fontsize=13.5, weight="bold", color=C_FLASH) T = 8 gx0, gy0, cell = 1.05, 3.1, 0.46 for i in range(T): for j in range(T): if i + 1 <= j: # 因果掩码下整块跳过(先判,优先级最高) fc, ec = C_SKIP, "#c3ccd8" elif j == 3: # 当前正在处理的 K/V 块列 fc, ec = "#cdeadb", C_SRAM elif i == 5: # 当前 Q 块行 fc, ec = "#dfeaf6", "#94a3b8" else: fc, ec = "#f7fafc", "#dde3ea" ax2.add_patch(Rectangle( (gx0 + j * cell, gy0 + (T - 1 - i) * cell), cell * 0.90, cell * 0.90, facecolor=fc, edgecolor=ec, linewidth=1.1)) gtop = gy0 + T * cell ax2.text(gx0 - 0.42, gy0 + T * cell / 2, "Q 块逐行流过", fontsize=9.5, color="#475569", ha="center", va="center", rotation=90) ax2.add_patch(FancyBboxPatch( (gx0 + 3 * cell - 0.06, gy0 - 0.06), cell * 1.02, T * cell + 0.12, boxstyle="round,pad=0.03,rounding_size=0.1", linewidth=1.6, edgecolor=C_SRAM, facecolor="none", linestyle="--")) ax2.text(gx0 + 3.5 * cell, gy0 - 0.42, "K_j, V_j 常驻 SRAM", ha="center", fontsize=9.5, color=C_SRAM, weight="bold") ax2.text(gx0 + T * cell + 0.45, gtop - 0.35, "一块 = [B_r, B_c]\n= [64, 64]", fontsize=9.5, color="#475569", va="top", linespacing=1.5) ax2.text(gx0 + T * cell + 0.45, gtop - 1.55, "灰色块:因果掩码下\n整块跳过,连算都不算", fontsize=9, color="#94a3b8", va="top", linespacing=1.5) ax2.text(gx0 + T * cell + 0.45, gtop - 2.95, "片上空间受限:\n驻留 K/V 与 Q/O 块,\n还要容纳分数和统计量", fontsize=9, color=C_SRAM, va="top", linespacing=1.5) by = 1.85 ax2.text(0.15, by + 0.52, "跨 K/V 块一直复用的三个量:", fontsize=10.5, color=INK, weight="bold") for k, (label, color) in enumerate([ ("O_i 输出累加器 [N, d]", C_FLASH), ("l_i 分母 [N]", "#7fb3d5"), ("m_i 运行最大值 [N]", "#c9a227")]): ax2.add_patch(Rectangle((0.15, by - k * 0.46), 4.6, 0.28, facecolor=color, edgecolor="none")) ax2.text(4.9, by - k * 0.46 + 0.14, label, fontsize=9.5, color="#475569", va="center") ax2.text(0.15, 0.12, "每来一块就 rescale 一次(乘 e^{m旧−m新}),最后 O/l 收尾", fontsize=9.5, color="#64748b") _caption(fig, "图 1:切的是「K/V 块 × Q 块」这两层循环,不是把注意力切开。" "左边突出两份二次中间量,右边只维护线性的累加量(固定 d)。") fig.savefig(OUT / "tiling.png", facecolor="#fbfcfe") plt.close(fig) print(" figures/tiling.png") # ──────────────────────────────────────────────────────────── # 图 2:两本账随 N 怎么长 # ──────────────────────────────────────────────────────────── def fig_curves(): d, b = 64, 2 B_r, B_c = block_sizes(d) Ns = np.array([1024, 2048, 4096, 8192, 16384, 32768]) naive_io = np.array([hbm_elems_naive(N, d) * b for N in Ns]) / MIB flash_io = np.array([hbm_elems_flash(N, d, B_r, B_c) * b for N in Ns]) / MIB fig, axes = plt.subplots(1, 2, figsize=(13.2, 5.1), facecolor="#fbfcfe") fig.subplots_adjust(bottom=0.20, top=0.86, wspace=0.25) ax = axes[0] ax.loglog(Ns, naive_io, "o-", color=C_NAIVE, lw=2.2, ms=6, label="朴素实现(∝N²)") ax.loglog(Ns, flash_io, "s-", color=C_FLASH, lw=2.2, ms=6, label="FlashAttention(∝N²/M)") ax.annotate(f"{naive_io[2] / flash_io[2]:.1f}×", xy=(Ns[2], naive_io[2]), xytext=(Ns[2] * 1.6, naive_io[2] * 3.0), fontsize=11, color=C_NAIVE, weight="bold", arrowprops=dict(arrowstyle="->", color=C_NAIVE, lw=1.3)) ax.legend(fontsize=10, frameon=False, loc="upper left") ax.grid(True, which="both", color="#eef2f7", lw=0.9) _style(ax, "HBM 读写量(单头 d=64, fp16)", "序列长度 N", "MiB") ax = axes[1] d2, H, B = 3072, 24, 1 lay = 32 naive_act = np.array( [activation_bytes(N, d2, H, B)[0] * lay for N in Ns]) / GIB flash_act = np.array( [activation_bytes(N, d2, H, B)[1] * lay for N in Ns]) / GIB ax.loglog(Ns, naive_act, "o-", color=C_NAIVE, lw=2.2, ms=6, label="朴素实现(∝N²)") ax.loglog(Ns, flash_act, "s-", color=C_FLASH, lw=2.2, ms=6, label="FlashAttention(∝N)") ax.axhline(80, color="#64748b", ls="--", lw=1.4) ax.text(Ns[0] * 1.15, 92, "示例预算:80 GiB(非硬件规格)", fontsize=9.5, color="#64748b") ax.annotate("76.5 GiB", xy=(4096, naive_act[2]), xytext=(4096 * 1.8, naive_act[2] * 2.2), fontsize=10.5, color=C_NAIVE, weight="bold", arrowprops=dict(arrowstyle="->", color=C_NAIVE, lw=1.3)) ax.annotate("4.5 GiB", xy=(4096, flash_act[2]), xytext=(4096 * 0.55, flash_act[2] * 6.0), fontsize=10.5, color=C_FLASH, weight="bold", arrowprops=dict(arrowstyle="->", color=C_FLASH, lw=1.3)) ax.legend(fontsize=10, frameon=False, loc="upper left") ax.grid(True, which="both", color="#eef2f7", lw=0.9) _style(ax, "教学激活存储模型(简化视频 DiT,32 层)", "序列长度 N", "GiB") _caption(fig, "图 2:两张都是对数轴,斜率就是复杂度阶数。" "左图的大 N 主导项均为 N²,有限 N 时比例受线性项和取整影响;" "右图斜率不同(N² 对 N),N 越大差距越离谱。") fig.savefig(OUT / "io_curve.png", facecolor="#fbfcfe") plt.close(fig) print(" figures/io_curve.png") # ──────────────────────────────────────────────────────────── # 图 3:算术强度 vs 机器平衡点 # ──────────────────────────────────────────────────────────── def fig_intensity(): d, b = 64, 2 B_r, B_c = block_sizes(d) Ns = np.array([1024, 2048, 4096, 8192, 16384, 32768]) naive_i = np.array([flops_attention(N, d) / (hbm_elems_naive(N, d) * b) for N in Ns]) flash_i = np.array([flops_attention(N, d) / (hbm_elems_flash(N, d, B_r, B_c) * b) for N in Ns]) balance = FLOPS_PEAK / HBM_BW fig, ax = plt.subplots(figsize=(11.4, 5.0), facecolor="#fbfcfe") fig.subplots_adjust(bottom=0.22, top=0.86) ax.axhspan(0, balance, color="#fdf1f2", zorder=0) ax.axhline(balance, color="#64748b", ls="--", lw=1.6, zorder=3) ax.text(Ns[-1] * 2.1, balance * 0.70, f"机器平衡点 {balance:.0f} FLOP/byte\n(峰值算力 ÷ HBM 带宽)", fontsize=10, color="#475569", va="center", linespacing=1.5) ax.text(Ns[0] * 1.02, balance * 0.40, "线下:本 roofline 模型由 HBM 项主导(实际瓶颈需测量)", fontsize=10, color=C_NAIVE) ax.text(Ns[0] * 1.02, balance * 1.32, "线上:本模型计算项主导", fontsize=10, color=C_FLASH) ax.semilogx(Ns, naive_i, "o-", color=C_NAIVE, lw=2.4, ms=7, label="朴素实现", zorder=4) ax.semilogx(Ns, flash_i, "s-", color=C_FLASH, lw=2.4, ms=7, label="FlashAttention", zorder=4) for x, v in zip(Ns, naive_i): ax.annotate(f"{v:.0f}", (x, v), textcoords="offset points", xytext=(0, -16), ha="center", fontsize=9, color=C_NAIVE) for x, v in zip(Ns, flash_i): ax.annotate(f"{v:.0f}", (x, v), textcoords="offset points", xytext=(0, 9), ha="center", fontsize=9, color=C_FLASH) ax.set_xticks(Ns) ax.set_xticklabels([f"{n // 1024}K" for n in Ns]) ax.set_xlim(Ns[0] * 0.75, Ns[-1] * 4.5) ax.set_ylim(0, balance * 1.85) ax.legend(fontsize=10.5, frameon=False, loc="center right") ax.grid(axis="y", color="#eef2f7", lw=0.9) _style(ax, "算术强度:每搬 1 字节能做多少次运算(单头 d=64, fp16)", "序列长度 N", "FLOP / byte") _caption(fig, "图 3:朴素实现的强度几乎不随 N 变(一直贴在 32 附近)," "在此模型中由 HBM 项主导;FlashAttention 的估计抬过平衡点," "并不单独证明实际 kernel 已充分利用计算单元。") fig.savefig(OUT / "intensity.png", facecolor="#fbfcfe") plt.close(fig) print(" figures/intensity.png") if __name__ == "__main__": print("生成配图:") fig_tiling() fig_curves() fig_intensity()
2026年09月25日
1 阅读
0 评论
0 点赞
2026-09-25
AIGC 基本功|性能建模与 Profiling:算力、带宽与显存账本-Roofline
性能建模与 Profiling:算力、带宽与显存账本 所属方向:推理加速 | 难度:进阶 | 前置知识:混合精度与数值稳定性、自注意力机制的计算与显存账本 关键词:性能建模、Profiling、Roofline、算术强度、memory bandwidth、FLOPs、MFU、PyTorch Profiler、Nsight 01. 为什么需要它 先看一组在同一台笔记本(Apple silicon,numpy fp32)上此前实测出来的数字(第 04 节另给本次复核快照),出自文末附录的 machine_probe.py,可以自己复现: GEMM : F = 0.998 GFLOP, D = 8.1 MB, 实测 0.498 ms(2003 GFLOP/s) 逐元素 : F = 1.000 GFLOP, D = 12000 MB, 实测 189.649 ms(5 GFLOP/s) → FLOPs 基本相同,耗时差 381 倍。 两段计算量的浮点运算次数几乎一样(都约 1 GFLOP),耗时差了 381 倍。如果拿「FLOPs 少的算子更快」这类直觉去做优化决策,在这个数字面前是反的:这里逐元素算子与 GEMM 的 FLOPs 相同,却因为每次都要把数据从内存搬进搬出,被带宽死死卡住。 再比如一个真实场景的优化评审:有人提议「把 LayerNorm 内部换成 fp8 计算单元重写,算力能翻几倍」。查一下账(第 03 节会算):LayerNorm 的算术强度只有约 1.5 FLOP/Byte(下文不含 beta 偏置的简化 LayerNorm),远低于 A100 的 ridge point 200.6,是典型的带宽受限算子——给它换更快的算力,在只改变算力峰值、访存不变的模型中加速比是 1.00,一分钱收益都没有。反过来,同是 fp8,把它用在 decode 阶段的大权重 GEMM 上却有收益,但收益来自权重字节减半,不是算力翻倍。同一笔投资,用在哪类算子上,结论完全相反。 这篇的目的就是把这套「先算账再动手」的方法补齐:两个账本——速度账(算力 vs 带宽)和容量账(显存四项)——加上一套实测手段(profiler)。量化、算子融合、KV cache 管理、并行切分,所有推理优化节点的收益判断都站在这篇的地基上。 02. 最小可用理解 三句话讲完核心: 任何算子的耗时有一个下界,由两种资源里更慢的那个决定:$T \ge \max(F/P_{\text{peak}},\ D/\beta)$。$F$ 是算子要做的浮点运算数,$P_{\text{peak}}$ 是硬件与该指令类型匹配的峰值算力;$D$ 是算子要搬运的字节数,$\beta$ 是带宽。算得再快也快不过「数据没到」。 算术强度 $I = F/D$ 决定卡在哪:与 ridge point $I^{\ast} = P_{\text{peak}}/\beta$ 比较,$I < I^{\ast}$ 是带宽受限,优化方向是少搬字节(融合、量化、FlashAttention);$I > I^{\ast}$ 是算力受限,优化方向才是少算(更好算法、更低精度计算单元)。 显存容量是另一本独立的账:权重 + KV cache + 激活 + 额外开销四项加总,才决定一张卡能塞多少并发。「权重放得下就能跑」只覆盖了四项里的第一项。 03. 数学推导 3.1 时间下界为什么取 max 一个算子要做 $F$ 个浮点运算(FLOP,口径:一次乘加记 2 个 FLOP),硬件每秒最多做 $P_{\text{peak}}$ 个——就算计算单元一刻不停,也至少要 $F/P_{\text{peak}}$ 秒。同理,算子要把 $D$ 字节的数据在内存和计算单元之间搬个来回,总线每秒最多搬 $\beta$ 字节——至少要 $D/\beta$ 秒。这两件事用的是不同资源,理想情况下可以完全重叠(算上一批数据的同时搬下一批),所以总时间的下界是两者取 max: $$T \ge \max\left(\frac{F}{P_{\text{peak}}},\ \frac{D}{\beta}\right)$$ 注意这是下界:真实 kernel 还有 kernel launch、同步、缓存未命中、TLB miss 等额外开销,实测只会更慢。后面的实验用经验标定值和估算流量代入,此时算出的只是模型估计;计时波动或缓存会产生超过 100% 的达成率,并不违反物理下界。差距也不能全算成可消除的优化空间。 3.2 算术强度与 ridge point 定义算术强度: $$I = \frac{F}{D}$$ 物理含义:每从内存搬 1 字节数据,能换来多少次浮点运算。它取决于实现与所选存储层级;实际缓存命中和重复加载又与硬件有关,不能视为完全与机器无关。再定义机器的 ridge point: $$I^{\ast} = \frac{P_{\text{peak}}}{\beta}$$ 物理含义:这台机器「算」和「搬」一样快的分界点,单位都是 FLOP/Byte,所以可以比。把 $I$ 与 $I^{\ast}$ 代回 3.1 的下界: $I < I^{\ast}$(带宽受限):$D/\beta$ 那一项更大,$T \approx D/\beta$,模型性能上界 $F/T \le I \cdot \beta$——性能与算力峰值无关,只跟 $I$ 成正比,这就是 roofline 图上那条斜线; $I \ge I^{\ast}$(算力受限):$T \approx F/P_{\text{peak}}$,性能封顶在 $P_{\text{peak}}$,这就是平顶。 $P_{\text{peak}}$ 用哪个口径要非常小心:A100 SXM 的 BF16 dense 是 312 TFLOP/s,A100 不支持原生 FP8 Tensor Core;H100 SXM 的 BF16 dense 约 989、FP8 dense 约 1979 TFLOP/s,差 6.3 倍——口径选错,受限类型的判断直接反掉(第 06 节细说)。 3.3 给几类算子记账 GEMM:$A[M,K] \times B[K,N] \to C[M,N]$。每个输出元素要做 $K$ 次乘加,共 $M N K$ 次,乘加各记一次: $$F_{\text{gemm}} = 2MNK,\qquad D_{\text{gemm}} = b\,(MK + KN + MN)$$ $b$ 是每元素字节数(bf16 取 2)。读 $A$、读 $B$、写 $C$ 各一遍,统计量这类小东西忽略。$M$ 越大,权重 $B[K,N]$ 被摊得越薄,$I$ 越高——这解释了为什么大 batch 的 GEMM 是算力受限、小 batch 的 GEMM 是带宽受限。 逐元素算子:$n$ 个元素各做 1 次运算,读 2 份写 1 份: $$F = n,\qquad D = 3 b n,\qquad I = \frac{1}{3b}$$ $I$ 是常数(fp32 下约 0.083,bf16 下约 0.17),与 $n$ 无关。在固定 dtype 与访存模型下,增大 n 本身不会提高 I,因而不会像增加 GEMM 的复用维度那样跨越 roofline 分界。 softmax 的读写口径:普通三阶段实现约 3 读 2 写。若整行能驻留片上存储,融合实现可约 1 读 1 写,得到 2.5 倍的理想流量比;独立的 online normalizer 通常先流式求归一化量,再重读输入写出概率,为 2 读 1 写,不能混为一谈。本文用约 $5nd$ 作为普通 softmax 的操作计数;max、exp 与除法并非都能跑在 Tensor Core 上,online 递推还会增加标量运算。可对照 在线归一化论文与 Triton 整行融合示例的不同前提。 单头注意力($N$ 是 token 数,$D_{h}$ 是每头维度):两次 $N \times N \times D_{h}$ 的矩阵乘加 softmax: $$F_{\text{attn}} = 4N^{2}D_{h} + 5N^{2}$$ 朴素实现的分数矩阵 $S$ 和概率矩阵 $P$ 都要落显存,二次项访存约 $4N^{2}$;FlashAttention 让 $N^{2}$ 只留在片上 SRAM,若仅统计每份输入读取一次与输出写回一次,可得到不可避免的数据流量下界;真实分块内核会反复读入 K/V 或 Q 等数据: $$D_{\text{朴素}} = b\,(4N^{2} + 4ND_{h}),\qquad D_{\text{flash,min}} = b \cdot 4ND_{h}$$ 代入 $N=4096$、$D_{h}=128$、bf16(脚本 roofline_model.py 实算):朴素 $D = 138.41\ \text{MB}$、$I = 62.7$,落在 A100($I^{\ast} = 200.6$)的带宽受限区;理想最低流量 $D_{\mathrm{flash,min}} = 4.19\ \text{MB}$、$I = 2068$,跳进算力受限区。FLOPs 一动没动(×1.000),roofline 时间下界从 0.089 ms 降到 0.028 ms(×3.2),这是按最低流量计算的理想示例,33 倍不是实际 FlashAttention 访存或速度的测量;有限片上存储下还需分块 IO 模型。FlashAttention 那篇的完整递推在知识树的下一节点展开,这里先用 roofline 把它的收益定位清楚。 图 1:这张图要看三样——散点的横坐标是各算子的算术强度 $I$,点越靠右越「算得过来」;按 I 与 I 的相对位置给点作模型分类;点到上界的距离并不能单独证明真实瓶颈;FlashAttention 会改变访存量和横坐标,无法直接把这段垂直差距视为它的加速收益。* 3.4 优化收益的上限 把算力翻倍($P_{\text{peak}} \to 2P_{\text{peak}}$),加速比是: $$S_{\text{算力}} = \frac{\max(F/P_{\text{peak}},\ D/\beta)}{\max(F/(2P_{\text{peak}}),\ D/\beta)}$$ 带宽受限时分子分母都是 $D/\beta$,$S_{\text{算力}} = 1$:白花钱。算力受限时加速至多为 2;当 $1<I/I^{\ast}<2$ 时,算力翻倍会遇到带宽上限,加速为 $I/I^{\ast}$。带宽翻倍对称地反一次。roofline_model.py 第四节把第 3.3 节的每个算子都代了一遍(A100 口径): 算子 I/I* 算力×2 的加速 带宽×2 的加速 逐元素 add [16M] 0.001 1.00x 2.00x LayerNorm [4096, 3072] 0.007 1.00x 2.00x softmax 朴素三遍 [4096, 4096] 0.002 1.00x 2.00x GEMM 512x4096x4096 2.041 2.00x 1.00x attention 朴素 [N=4096,D=128] 0.312 1.00x 2.00x attention ideal-min [N=4096,D=128] 10.307 2.00x 1.00x 这张表就是「投资之前先看图 2」的数字版:你的 kernel 在分界线哪一侧,决定哪类投资是零收益。 图 2:这张图要看什么——横轴 $I/I^{\ast}=1$ 那条竖线就是分界线:线左边算力翻倍的理想收益为 1,带宽翻倍收益在 1 到 2 之间;只有 I/I≤0.5 时完整得到 2 倍。线右边对称,I/I≥2 才完整得到算力翻倍的 2 倍。投入硬件或投入算子融合之前,先看自己在哪一侧。 3.5 显存的容量账 速度账之外另有一本容量账。下面以未分块 prefill、同时处理 BS 个 token 为例,显存分四项: $$M_{\text{total}} = N_{\text{params}} \cdot b_{w} + 2BSLd_{\text{kv}}b + BS \cdot a + \rho \cdot M_{\text{sum}}$$ 逐项说物理含义:$N_{\text{params}}$ 是参数量、$b_{w}$ 是每参数字节数(fp16 为 2)——这一项与并发无关,是常数;$B$ 是并发序列数、$S$ 是序列长度、$L$ 是层数、$d_{\text{kv}}$ 是每层 KV 总维度(GQA 模型用实际的 KV head 数乘头维度)——KV cache 每个 token 每层都要存一份 K 和一份 V,所以是 $2BSLd_{\text{kv}}b$ 字节,随并发线性增长;$a$ 是每 token 的激活峰值(推理不保留整层中间结果,但当前层十几份临时张量要同时活着,$d_{\text{model}}=4096$、fp16 时取约 128 KiB/token 是经验值,随实现差距很大);$\rho$ 是额外开销率,$M_{\text{sum}}$ 是前三项之和。逐 token decode 时通常只有 B 个活跃 token,应将激活项 BS·a 改成 B·a;KV cache 仍随 BS 增长。分块 prefill 则按实际活跃 chunk 计。第 04 节把未分块 prefill 的教学账本代进一个 7B 模型。 04. 代码实现 三个脚本全部只用 numpy,因为 roofline 的方法论不依赖 GPU:同一台机器、同一套口径,把「峰值」和「落点」都实测出来,预测和实测的差距才看得见。本次复核实测峰值:$\beta = 71.8\ \text{GB/s}$、$P = 1635\ \text{GFLOP/s}$、$I^{\ast} = 22.8\ \text{FLOP/Byte}$(注意这是「numpy 能摸到的上限」,不是芯片标称值——方法论可比的前提是口径一致)。 4.1 标定两个峰值 带宽用两输入向量加法模式标定(不是含乘法的 STREAM triad):用足够大的工作集降低缓存影响;128 MB 是否超过目标机器缓存仍需核对,测得的是该访问模式的有效带宽。 def measure_bandwidth(n: int = 32_000_000, repeat: int = 8) -> float: """z = x + y:读 2n、写 n,共 3n 个 fp32 元素。""" x = np.ones(n, dtype=np.float32) y = np.ones(n, dtype=np.float32) z = np.empty(n, dtype=np.float32) dt = _best(lambda: np.add(x, y, out=z), repeat) return 3 * n * FP32 / dt # Byte/s,_best 取 repeat 次最快 算力用足够大的方阵乘标定(4096 的方阵乘访存被摊薄,$F = 2n^{3}$)。实跑输出: 本机标定(arm64 / numpy 2.1.3 / fp32,2026-10-02 02:21:20) 实测可达带宽 beta = 71.78 GB/s 实测可达算力 P = 1635.04 GFLOP/s ridge point I* = 22.8 FLOP/Byte 4.2 把真实算子打上 roofline 关键测试对象是 matmul-softmax-matmul 微基准;此处省略 1/√D 缩放,不是完整的生产 attention。下面这七行就是「账本」本身——每一行右边标了它读写了几个 $N^{2}$ 量级的遍数,$D$ 就是这么数出来的,不是拍脑袋: def attention_naive(): np.matmul(q, kt, out=S) # 写 S 1 遍 np.max(S, axis=-1, keepdims=True, out=rowmax) # 读 S 2 遍 np.subtract(S, rowmax, out=S) # 读写 S 4 遍 np.exp(S, out=S) # 读写 S 6 遍 np.sum(S, axis=-1, keepdims=True, out=rowsum) # 读 S 7 遍 np.divide(S, rowsum, out=S) # 读写 S 9 遍 np.matmul(S, v, out=O) # 读 S 写 O 10 遍 实跑结果(本机口径,$I^{\ast} = 22.8$): 算子 I 模型分类 T估计 T实测 达成率 实测算力 逐元素 add(1 遍) 0.08 带宽 5.35ms 5.07ms 106% 6.3 GF/s add+relu 两遍(预分配) 0.10 带宽 8.92ms 16.76ms 53% 3.8 GF/s add+relu 写成一行(有中间数组) 0.10 带宽 8.92ms 26.80ms 33% 2.4 GF/s LayerNorm [8192,4096] 预分配 0.15 带宽 18.70ms 32.56ms 57% 6.2 GF/s attention 朴素 [N=2048,D=128] 12.61 带宽 2.40ms 11.41ms 21% 190.1 GF/s GEMM 512x4096x4096 204.80 算力 10.51ms 16.18ms 65% 1061.5 GF/s GEMM 4096x4096x4096 682.67 算力 84.06ms 84.82ms 99% 1620.3 GF/s 逐行解读,三种达成率各说明一件事: 逐元素 add 106%:与标定时相同访问模式,已接近本次有效带宽。超过 100% 来自经验标定与实际计时的差异,不能理解成超过硬件物理峰值。 GEMM 大矩阵 99%:已接近本机同类 GEMM 标定值;这不能证明其他实现没有改进空间。 朴素 attention 21%:实测约为模型估计时间的 4.8 倍,可能涉及缓存、指令吞吐、线程和内核调度;仅凭总时间不能分离原因,也不能直接解释成多搬了几倍字节。 图 1 使用完整脚本中的 8 个算子(正文节选了 7 行)。点到 roofline 的距离表示相对模型上界的性能差距;它本身不能区分额外访存、指令开销或同步等原因。 4.3 显存账本与 batch 拐点 未分块 prefill 的容量教学账代一个 7B 模型(32 层、$d_{\text{kv}} = 4096$、fp16): 模型:7B,fp16 权重 = 13.0 GiB,32 层,d_kv = 4096 KV cache 单价:512 KiB / token (一条 4096 长的序列 = 2.00 GiB) batch 权重 KV cache 激活 额外开销 合计 假设 78 GiB 可用预算 1 13.0G 2.00G 0.50G 4.66G 20.20G 装得下 8 13.0G 16.00G 4.00G 9.91G 42.95G 装得下 16 13.0G 32.00G 8.00G 15.91G 68.95G 装得下 24 13.0G 48.00G 12.00G 21.91G 94.95G OOM 32 13.0G 64.00G 16.00G 27.91G 120.95G OOM 两个读数:权重是常数,batch 再大都是 13.0 GiB;KV cache 按 512 KiB/token 线性涨,batch 16 时(32 GiB)已经是权重的 2.5 倍。此处 30% 是人为设定的额外开销/前三项总量比率,不是实测额外开销率;它也不同于 PagedAttention 论文中 KV cache 已分配容量的浪费比例。同一张卡只把本教学模型的额外开销率改成 4%,batch 24 从 94.95 GiB 降到 75.96 GiB,从 OOM 变装得下:不换卡、不改模型,这个假设账本跨过一个 batch 档位;真实分页收益需只对可优化的 KV 分配项建模,并检查实际可用显存。 容量账决定了「能开多大」,速度账决定「该开多大」。固定一份 64 MiB 的权重矩阵,扫 batch(每个 batch 位置喂一个 token),本机实测: B F (GFLOP) D (MB) I 模型分类 延迟 ms 吞吐 K/s 1 0.034 67.14 0.50 带宽 1.718 0.6 2 0.067 67.17 1.00 带宽 8.232 0.2 4 0.134 67.24 2.00 带宽 8.161 0.5 8 0.268 67.37 3.98 带宽 4.477 1.8 16 0.537 67.63 7.94 带宽 4.436 3.6 32 1.074 68.16 15.75 带宽 4.461 7.2 64 2.147 69.21 31.03 算力 4.865 13.2 128 4.295 71.30 60.24 算力 6.362 20.1 256 8.590 75.50 113.78 算力 8.612 29.7 延迟:B=1 时 1.718 ms,B=256 时 8.612 ms,涨 5.0 倍;吞吐涨 51.1 倍。理想带宽模型中,加大 batch 可摊薄权重读取;本机表格并不呈严格平坦段。B=32 到 64 之间越过模型分界 $I^{\ast}=22.8$,这不保证实测曲线在同一处出现锐利拐点。B=256 时 64 MiB 的权重读取被摊到每样本 0.25 MiB。 B=2 的延迟是 B=8 的 1.84 倍,可能涉及 BLAS 内核选择、线程调度或计时波动;仅凭这张时间表不能确证原因。roofline 看不见这类事情,定 batch 必须结合实测。 图 3:这张图要看什么——左边是实测延迟,右边是实测吞吐,曲线不保证出现理想平坦段;竖虚线是本次标定值代入模型后的分类交界;B=2 处那个尖是 roofline 模型解释不了的库行为,标出来是为了让你对「模型给方向、实测给结论」有体感。 图 4:这张图要看什么——左图是显存四项的堆叠:蓝色权重是常数,红色 KV cache 随并发线性涨,越过黑色可用线就是 OOM;右图是只改一个参数(额外开销率 30%→4%)的效果,batch 24 从 OOM 变装得下。 05. 工业级实现对照 最小实现是方法论,生产里测量走的是另一套工具,但问的是同一组问题。 PyTorch Profiler(pytorch/pytorch · torch/profiler/profiler.py · profile,以 2026-09 的实现为准)是最常用的第一站: from torch.profiler import profile, ProfilerActivity with profile(activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA], record_shapes=True, profile_memory=True, with_flops=True) as prof: model(x) print(prof.key_averages().table(sort_by="cuda_time_total", row_limit=15)) prof.export_chrome_trace("trace.json") 和本文方法的对应关系:key_averages().table 按算子事件键聚合 CPU/CUDA 耗时,而非直接按 GPU kernel 名称聚合,回答「时间花在哪」;with_flops=True 只为支持的算子(如矩阵乘与二维卷积)估算 FLOPs,不能视为全模型所有操作的完整计数,除以耗时和相应算力峰值得到该算子的利用率估计,并不是通常定义的模型级 MFU(Model FLOPs Utilization,模型理论 FLOPs ÷ 墙钟时间 ÷ 硬件峰值算力);profile_memory=True 抓张量级显存分配,对应容量账的激活项。训练报告里常见的整体 MFU 是同一口径的粗化:拿模型理论 FLOPs 除以墙钟时间和卡数峰值,这个数在 decode 型负载里天然上不去,原因见第 08 节第 4 条。 Nsight 全家桶是更细的一层:Nsight Systems(nsys)看时间线——kernel 之间的空隙、同步等待、通信重叠,对应「下界假设完美重叠」不成立的部分;Nsight Compute(ncu)看单个 kernel,它的 Speed Of Light 面板直接给出 Compute Throughput 和 Memory Throughput 两个百分比——那就是 roofline 的粗版:应结合指令流水线、occupancy、缓存层级和 stall 原因判断,不能仅比较两个百分比便确定瓶颈。实操顺序通常是:nsys 找到热点和空隙,ncu 对热点 kernel 看 SOL 定受限类型,再决定投算力还是投带宽。 和框架选择的关系:SDPA 会根据输入与硬件分派后端,不能断言所有 diffusers/transformers 都默认使用 FlashAttention。本文 2068 FLOP/byte 来自每份 Q/K/V 只读一次的理想下界;真实 FlashAttention 还需考虑分块重读。小 batch GEMM 的权重读取成本通常很高,但具体量化收益与核实现、反量化和 batch 都有关。 06. 代价与边界 roofline 是模型,模型有假设。四条主要假设和不成立时的样子: 假设计算与访存完全重叠。真实 kernel 在算和搬之间来回切换,还有 launch 和同步的开销。小 kernel(本机实测 B=1 约 1.7 ms 的场景)里这些固定开销占比不小,模型系统性偏乐观。 假设峰值算力是一个数。实际有 fp16/bf16/fp8/稀疏好几档,差 6 倍以上;按 BF16 dense 口径,H100 SXM 相比 A100 40GB SXM 的 ridge point 从约 201 升到 295——换新卡后更多 kernel 会落进带宽受限区,拿旧卡的直觉做判断会错。 假设 $D$ 与缓存无关。账本里的 $D$ 按落盘遍数数,但缓存命中会让有效带宽远大于 DRAM 标称值。锚点论文 Hierarchical Roofline(arXiv:2009.05257)就是把单一 roofline 扩展成每级缓存一条,用于定位数据移动发生在哪一层。 假设算子孤立。decode 阶段 GEMM 的「权重」每层都要重新读一遍,全局账(整个模型、整个请求)和单算子账结论可能不同;通信算子(allreduce)的账本里延迟和消息数占大头,照搬本文公式会算错。 什么时候不用它:动态 shape、算子间强耦合(融合边界在变)、通信密集的分布式场景——这些先上 profiler 看时间线,roofline 只对「单 kernel、口径清晰」的问题给下界。坦诚标注:本文所有「实测」都来自一台笔记本的 numpy,数字本身不可迁移,理想分段趋势也可能被缓存、内核选择和调度打破。 07. 经典论文脉络 Roofline: An Insightful Visual Performance Model(Williams et al., CACM 2009,未挂 arXiv):提出算术强度与 ridge point,一根折线把「算力受限/带宽受限」变成可判定的题。一切性能建模的原点。 Hierarchical Roofline Performance Analysis for Deep Learning Applications(Yang et al., 2020):把 roofline 按缓存层级展开,回答「多搬的字节发生在哪一级存储」,是本文锚点论文,也补了单一 roofline 在深度学习负载上最大的盲区。 FlashAttention(Dao et al., 2022):IO-aware 的代表作——FLOPs 一动不动,靠 tiling + online softmax 减少 N×N 中间量落地及相应 IO,把注意力从带宽受限拉进算力受限。roofline 视角下「优化带宽」的教科书案例。 Mixed Precision Training(Micikevicius et al., 2017):换 dtype 同时改两本账——$b$ 变小省字节,算力单元换挡提峰值。哪半边有收益取决于受限类型(已发长文专门讲数值稳定那半)。 Efficient Memory Management with PagedAttention(Kwon et al., 2023):在论文比较的服务负载中显著减少 KV cache 分配浪费;这些比例不能直接乘到权重、激活和所有显存上,虚拟内存的分页思想搬进显存管理。知识树里 KV cache 一篇的主角。 五篇连起来是一条线:先有判定工具(roofline),再按受限类型各给一把钥匙——算力侧(混合精度)、带宽侧(FlashAttention)、容量侧(PagedAttention)。 08. 常见误解 「FLOPs 少的算子更快」。本文开头的实测:同样约 1 GFLOP,GEMM 0.498 ms,逐元素 189.6 ms,差 381 倍。FLOPs 只在算力受限区才和耗时挂钩;带宽受限区里,决定耗时的是字节数。 「显存够放权重就能跑」。7B fp16 权重只要 13 GiB,但 batch 24 时 KV cache 48 GiB + 额外开销 21.9 GiB,假设有 78 GiB 可用预算的设备照样 OOM。容量账要四项加总,KV cache 那一项随并发线性涨,batch 16 就反超权重了。 「新卡算力翻倍,我的推理一定提速」。带宽受限的 kernel 加速比是 1.00(第 3.4 节表格里的 1.00x)。H100 SXM 的 FP8 dense 峰值约为 A100 BF16 的 6.3 倍,但实际换卡同时还改变带宽与内核;带宽受限时应估算字节数/实际带宽。权重量化会减少读取量,单独提高计算峰值则未必有用。 「MFU 低就是实现烂」。decode 阶段每个 token 过一遍全部权重,$I$ 天然低于 ridge point,MFU 高不了——这是负载形状决定的,不是代码烂。看 MFU 前先分清 典型 prefill 与小 batch decode 的负载形状;足够大的 decode batch、长上下文注意力或通信可能改变瓶颈。 「账本算出来的就是实际」。朴素 attention 模型估计 2.40 ms,实测 11.41 ms(差 4.8 倍);batch 扫描里 B=2 延迟是 B=8 的 1.84 倍(原因需进一步 profile)。模型给方向和上限,实测给结论——两个都要,缺一个都会做出错误决策。 09. 动手验证 三个都能在笔记本上跑(附录有完整代码): python machine_probe.py——预期:逐元素 add 达成率 ≈100%,GEMM 大矩阵 ≈100%,朴素 attention 明显低于 50%。如果你的机器上 attention 达成率反而很高,多半是缓存把 $N^{2}$ 矩阵装下了,把 N 调大一倍再看。 python memory_ledger.py——观察实际延迟/吞吐,不预设理想三段式;按 I 与 I* 标记模型分类交界(本机在 32→64 之间)。改 batch_sweep 里的 B 列表,看吞吐什么时候不再涨。 打开 roofline_model.py,仅把 A100 示例的 peak_flops 改成 1979e12、其余不变,重跑第五节(这是控制变量实验,不代表实际 H100,真实换卡还需更新带宽与 dtype)——预期:GEMM 拐点从 M∈(128, 256] 右移,更多算子被判为带宽受限。这一步会让你体感 ridge point 抬高对优化决策的影响。 10. 延伸阅读 按知识树的依赖关系,从这篇出发有三个方向: 往上游:自注意力机制的计算与显存账本——本文 3.3 节那笔注意力账的完整推导;混合精度与数值稳定性——换 dtype 的另一本账(数值范围与累加精度)。 往下游(受限类型各一把钥匙):FlashAttention(带宽侧)、KV Cache 与自回归视频生成(容量侧)、算子融合与 CUDA Graph(launch 开销侧,已排期)、量化(字节侧,已排期)。 目录:完整知识树见博客目录页 AIGC 基本功知识树,按依赖顺序排好了先修课。 继续阅读 FlashAttention 为什么不需要存下注意力矩阵:本文 3.3 节的最低流量示例省略了真实分块重读,那篇把 online softmax 的递推式一步步推出来。 附录:完整代码 09 节用到的脚本全文如下(machine_probe.py、roofline_model.py、memory_ledger.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 machine_probe.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """在同一台机器上把 roofline 的「峰值」和「落点」都实测出来。 只依赖 numpy。直接 `python machine_probe.py` 即可运行,约 20 秒。 为什么要自己测一遍 ------------------ roofline_model.py 里的 A100/H100 峰值是厂商标称值,读者手边不一定有 GPU。 这个脚本用同一套口径(F 记 2·MAC、D 记读写字节)在同一台机器上先标定 峰值算力与峰值带宽,再把真实算子打到这张 roofline 上——所以你能看到 「预测的下界」和「实测」差多少,而不只是相信规格表。 每个算子的 D 都按实现里实际的读写遍数来数: LayerNorm 预分配版 10 遍(mean / 减均值 / 平方 / 再 mean / 除 / 乘 gamma) attention 朴素预分配 10 遍(写 S / max / 减 / exp / sum / 除 / 读 S 写 O) 遍数不是拍脑袋来的,是照着下面每一行 numpy 数出来的。 输出的每个数字都来自 `time.perf_counter()` 的真实计时。 """ import json import platform import time from pathlib import Path import numpy as np FP32 = 4 # 本机跑 fp32,numpy 在 Apple 上走 Accelerate SNAPSHOT = Path(__file__).resolve().parent / "probe_results.json" def _best(fn, repeat: int): """跑 repeat 次取最快的一次:避开冷启动和调度抖动。""" fn() # 预热:第一次要分配 / 触发缺页 best = float("inf") for _ in range(repeat): t = time.perf_counter() fn() best = min(best, time.perf_counter() - t) return best # ── 一、标定可达带宽:二元 add 模式(读两份写一份)─────────────────── def measure_bandwidth(n: int = 32_000_000, repeat: int = 8) -> float: """返回实测可达带宽(Byte/s)。 每份 fp32 数组 128 MB;应核对目标机器缓存大小,这里测量此访问模式的有效带宽。 z = x + y:读 2n、写 n,共 3n 个元素。 """ x = np.ones(n, dtype=np.float32) y = np.ones(n, dtype=np.float32) z = np.empty(n, dtype=np.float32) dt = _best(lambda: np.add(x, y, out=z), repeat) return 3 * n * FP32 / dt # ── 二、标定峰值算力:大矩阵乘 ──────────────────────────────────── def measure_matmul_peak(n: int = 4096, repeat: int = 5) -> float: """返回实测可达算力(FLOP/s)。 4096 的方阵乘足够大,访存被摊薄,测出来的是算力上限。 FLOPs = 2·n³(口径同 roofline_model.py:一次乘加记 2 个浮点运算)。 """ rng = np.random.default_rng(0) a = rng.standard_normal((n, n), dtype=np.float32) b = rng.standard_normal((n, n), dtype=np.float32) c = np.empty((n, n), dtype=np.float32) dt = _best(lambda: np.matmul(a, b, out=c), repeat) return 2 * n ** 3 / dt def calibrate(verbose: bool = True): bw = measure_bandwidth() fl = measure_matmul_peak() peak = dict(peak_flops=fl, peak_bw=bw) if verbose: print("=" * 80) print(f"本机标定({platform.machine()},numpy {np.__version__},fp32)") print("=" * 80) print(f" 实测可达带宽 beta = {bw / 1e9:8.2f} GB/s") print(f" 实测可达算力 P = {fl / 1e9:8.2f} GFLOP/s") print(f" ridge point I* = {fl / bw:8.1f} FLOP/Byte") print(" 注:这是「用 numpy 能摸到」的上限,不是硬件标称值;") print(" 换 BLAS 后端、换 dtype、换线程数都会变。口径一致才有可比性。") return peak def build_ops(): """构造待实测的算子。每个元素 = (名字, F, D, 函数, repeat)。""" rng = np.random.default_rng(0) ops = [] # ── 逐元素:1 遍 vs 2 遍 vs「看起来融合了」 ── n = 32_000_000 a = rng.standard_normal(n).astype(np.float32) b = rng.standard_normal(n).astype(np.float32) o1 = np.empty(n, dtype=np.float32) tmp = np.empty(n, dtype=np.float32) def add_relu_two_pass(): # 5n 字节:tmp 读写各一次 np.add(a, b, out=tmp) np.maximum(tmp, 0, out=o1) def add_relu_one_line(): # 写法像融合了,numpy 照样分配中间数组 np.maximum(np.add(a, b), 0, out=o1) ops.append(("逐元素 add(1 遍)", 1.0 * n, FP32 * 3 * n, lambda: np.add(a, b, out=o1), 8)) ops.append(("add+relu 两遍(预分配)", 2.0 * n, FP32 * 5 * n, add_relu_two_pass, 8)) ops.append(("add+relu 写成一行(有中间数组)", 2.0 * n, FP32 * 5 * n, add_relu_one_line, 8)) # ── LayerNorm:预分配的 10 遍版 ── ln_n, ln_d = 8192, 4096 x = rng.standard_normal((ln_n, ln_d)).astype(np.float32) g = rng.standard_normal(ln_d).astype(np.float32) mu = np.empty((ln_n, 1), dtype=np.float32) var = np.empty((ln_n, 1), dtype=np.float32) xc = np.empty_like(x) xc2 = np.empty_like(x) def layernorm_buffered(): np.mean(x, axis=-1, keepdims=True, out=mu) # 读 x 1 遍 np.subtract(x, mu, out=xc) # 读 x 写 xc 3 遍 np.multiply(xc, xc, out=xc2) # 读 xc 写 xc2 5 遍 np.mean(xc2, axis=-1, keepdims=True, out=var) # 读 xc2 6 遍 np.sqrt(var + 1e-5, out=var) # 小量 np.divide(xc, var, out=xc) # 读 xc 写 xc 8 遍 np.multiply(xc, g, out=xc) # 读 xc 写 xc 10 遍 ops.append((f"LayerNorm [{ln_n},{ln_d}] 预分配", 6.0 * ln_n * ln_d, FP32 * 10 * ln_n * ln_d, layernorm_buffered, 5)) # ── 朴素注意力:S 落盘,10 遍 N² 访存 ── N, D = 2048, 128 q = rng.standard_normal((N, D)).astype(np.float32) k = rng.standard_normal((N, D)).astype(np.float32) v = rng.standard_normal((N, D)).astype(np.float32) kt = np.ascontiguousarray(k.T) S = np.empty((N, N), dtype=np.float32) O = np.empty((N, D), dtype=np.float32) rowmax = np.empty((N, 1), dtype=np.float32) rowsum = np.empty((N, 1), dtype=np.float32) def attention_naive(): # 本微基准省略 1/sqrt(D) 缩放,只测 matmul-softmax-matmul 的执行开销。 np.matmul(q, kt, out=S) # 写 S 1 遍 np.max(S, axis=-1, keepdims=True, out=rowmax) # 读 S 2 遍 np.subtract(S, rowmax, out=S) # 读写 S 4 遍 np.exp(S, out=S) # 读写 S 6 遍 np.sum(S, axis=-1, keepdims=True, out=rowsum) # 读 S 7 遍 np.divide(S, rowsum, out=S) # 读写 S 9 遍 np.matmul(S, v, out=O) # 读 S 写 O 10 遍 ops.append((f"attention 朴素 [N={N},D={D}]", 4.0 * N * N * D + 5.0 * N * N, FP32 * (10 * N * N + 4 * N * D), attention_naive, 5)) # ── GEMM 三个规模 ── for m in (512, 2048, 4096): ma = rng.standard_normal((m, 4096)).astype(np.float32) mb = rng.standard_normal((4096, 4096)).astype(np.float32) mc = np.empty((m, 4096), dtype=np.float32) ops.append((f"GEMM {m}x4096x4096", 2.0 * m * 4096 * 4096, FP32 * (m * 4096 + 4096 * 4096 + m * 4096), lambda ma=ma, mb=mb, mc=mc: np.matmul(ma, mb, out=mc), 5)) return ops def probe(peak: dict, verbose: bool = True): ops = build_ops() I_star = peak["peak_flops"] / peak["peak_bw"] if verbose: print("\n" + "=" * 80) print("三、真实算子打到这张 roofline 上:下界 vs 实测") print("=" * 80) hdr = (f"{'算子':<32s}{'I':>9s}{'受限':>8s}{'T下界':>10s}{'T实测':>10s}" f"{'达成率':>9s}{'实测算力':>14s}") print(hdr) print("-" * len(hdr)) rows = [] for name, F, D, fn, repeat in ops: dt = _best(fn, repeat) I = F / D bound = "算力" if I >= I_star else "带宽" t_pred = max(F / peak["peak_flops"], D / peak["peak_bw"]) eff = t_pred / dt rows.append(dict(name=name, I=I, bound=bound, t_pred=t_pred, t_real=dt, eff=eff, flops=F, bytes=D, gflops=F / dt / 1e9)) if verbose: print(f"{name:<32s}{I:>9.2f}{bound:>8s}{t_pred * 1e3:>8.2f}ms" f"{dt * 1e3:>8.2f}ms{eff * 100:>8.0f}%" f"{F / dt / 1e9:>11.1f} GF/s") if verbose: print("\n 达成率 = 下界 / 实测,回答的是「这个实现离 roofline 还有多远」。") print(" 低达成率可能来自指令、缓存、调度或同步;单凭此比值不能诊断额外访存。") return rows def demo_same_flops(): """同样 1 GFLOP 的算力,GEMM 和逐元素算子差多少时间。""" print("\n" + "=" * 80) print("四、同样的 FLOPs,两种算子差多少时间") print("=" * 80) rng = np.random.default_rng(3) target = 1.0e9 # GEMM:2·M·K·N = target,取 K=N=1024 M = int(target / (2 * 1024 * 1024)) ma = rng.standard_normal((M, 1024)).astype(np.float32) mb = rng.standard_normal((1024, 1024)).astype(np.float32) mc = np.empty((M, 1024), dtype=np.float32) t_gemm = _best(lambda: np.matmul(ma, mb, out=mc), 10) f_gemm = 2.0 * M * 1024 * 1024 # 逐元素:n 个元素各算 1 次,F = n n_elem = int(target) a = rng.standard_normal(n_elem).astype(np.float32) b = rng.standard_normal(n_elem).astype(np.float32) o = np.empty(n_elem, dtype=np.float32) t_elem = _best(lambda: np.add(a, b, out=o), 5) print(f" GEMM : F = {f_gemm / 1e9:.3f} GFLOP, " f"D = {FP32 * (M * 1024 + 1024 * 1024 + M * 1024) / 1e6:.1f} MB, " f"实测 {t_gemm * 1e3:.3f} ms({f_gemm / t_gemm / 1e9:.0f} GFLOP/s)") print(f" 逐元素 : F = {n_elem / 1e9:.3f} GFLOP, " f"D = {FP32 * 3 * n_elem / 1e6:.1f} MB, " f"实测 {t_elem * 1e3:.3f} ms({n_elem / t_elem / 1e9:.0f} GFLOP/s)") print(f" → FLOPs 基本相同,耗时差 {t_elem / t_gemm:.0f} 倍。") print(" 同样运算量不代表同样时间:这里逐元素算子的运算数并未") print(f" 减少,却因为 I 只有 {1 / (FP32 * 3):.2f} FLOP/Byte 而完全被带宽卡住。") def main(): peak = calibrate() rows = probe(peak) demo_same_flops() # 把这次标定存成快照:make_figures / memory_ledger 直接读它, # 保证「图上的数字」和「正文引用的数字」来自同一次运行。 SNAPSHOT.write_text(json.dumps( dict(peak=peak, rows=rows, note=f"{platform.machine()} / numpy {np.__version__} / fp32", timestamp=time.strftime("%Y-%m-%d %H:%M:%S")), ensure_ascii=False, indent=1)) print(f"\n标定快照已存到 {SNAPSHOT.name},make_figures / memory_ledger 会复用它。") print(f"这台机器的 roofline:P = {peak['peak_flops'] / 1e9:.0f} GFLOP/s, " f"beta = {peak['peak_bw'] / 1e9:.1f} GB/s, " f"I* = {peak['peak_flops'] / peak['peak_bw']:.1f} FLOP/Byte") if __name__ == "__main__": main() roofline_model.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """Roofline 模型的最小实现:把「算力账」和「带宽账」合成一条上界曲线。 只依赖 numpy(其实只是用来排版,算法本身是纯标量运算)。 直接 `python roofline_model.py` 即可运行,约 1 秒。 符号与文章第 03 节一一对应 -------------------------- F 浮点运算数(FLOP)。口径:一次乘加(MAC)记 2 FLOP D 访存字节数(Byte)。口径:读 + 写,缓存命中不重复计 I 算术强度 I = F / D,单位 FLOP/Byte P_peak 峰值算力(FLOP/s) beta 峰值带宽(Byte/s) I_star ridge point = P_peak / beta,I 低于它就是带宽受限 关键结论(代码算完会打印): T >= max(F / P_peak, D / beta) 把 F 降为一半或把 D 降为一半时的模型比值,见 optimize_headroom() """ import json from pathlib import Path def _local_spec(): """本机实测峰值:优先读 machine_probe.py 存下的标定快照。 快照不存在时退回兜底常数(先跑一次 machine_probe.py 更准)。 这样 SPECS 里的「本机实测」永远是同一次标定,不会和正文引用的数字分叉。 """ snap = Path(__file__).resolve().parent / "probe_results.json" if snap.exists(): p = json.loads(snap.read_text())["peak"] return dict(peak_flops=p["peak_flops"], peak_bw=p["peak_bw"]) return dict(peak_flops=1680e9, peak_bw=77e9) # ── 厂商标称峰值。来源:NVIDIA A100 / H100 SXM 白皮书,bf16 dense(不含稀疏) # 注意 H100 BF16 dense 算力增幅大于带宽增幅,ridge point 比 A100 高: # 带宽涨 2.15 倍、bf16 算力涨约 3.17 倍 → 更多 kernel 落在带宽受限区。 SPECS = { "A100-40GB SXM (bf16)": dict(peak_flops=312e12, peak_bw=1555e9), "H100-80GB SXM (bf16)": dict(peak_flops=989e12, peak_bw=3350e9), "H100-80GB SXM (fp8)": dict(peak_flops=1979e12, peak_bw=3350e9), # 下面这一条不是标称值,是 machine_probe.py 在同一台机器上实测出来的 "本机实测(见 machine_probe.py)": _local_spec(), } BF16 = 2 # bytes per element def roofline_time(F: float, D: float, peak_flops: float, peak_bw: float): """roofline 时间下界:算力时间与带宽时间取 max(假设两者完美重叠)。""" t_compute = F / peak_flops t_memory = D / peak_bw return max(t_compute, t_memory), t_compute, t_memory def bound_of(F: float, D: float, peak_flops: float, peak_bw: float) -> str: I = F / D I_star = peak_flops / peak_bw return "算力受限" if I >= I_star else "带宽受限" def optimize_headroom(F: float, D: float, peak_flops: float, peak_bw: float): """把运算量 F 或搬运量 D 减半,理想模型可快多少? 答案与算术强度有关:带宽受限时砍算力收益为 0,算力受限时最多 2 倍。 返回 (F 减半的加速比, D 减半的加速比),分别等价于只将峰值算力或带宽翻倍。 """ def speedup(F2, D2): t0, _, _ = roofline_time(F, D, peak_flops, peak_bw) t1, _, _ = roofline_time(F2, D2, peak_flops, peak_bw) return t0 / t1 return speedup(F / 2, D), speedup(F, D / 2) # ── 几类算子的 F / D 账本 ────────────────────────────────────────── def gemm(M: int, K: int, N: int, b: int = BF16): """矩阵乘 [M,K] @ [K,N] -> [M,N]。读 A、读 B、写 C。""" F = 2.0 * M * K * N D = b * (M * K + K * N + M * N) return F, D def elementwise(n: int, b: int = BF16, n_pass: int = 1, n_read: int = 1): """逐元素算子:读 n_read 份、写 1 份;n_pass 表示这样读写几轮。 二元运算(如 add)是 n_read=2:读两份输入、写一份输出。 """ F = 1.0 * n * n_pass # 每个元素算 1 次 D = b * (n_read + 1) * n * n_pass return F, D def layernorm(n: int, d: int, b: int = BF16): """简化 LayerNorm(无 beta):两次归约、减均值、平方、除标准差、乘 gamma,约 6 次/元素。 访存 = 读 x 一遍 + 写 y 一遍(统计量是 O(n) 的小量,忽略)。""" F = 6.0 * n * d D = b * 2 * n * d return F, D def softmax(n: int, d: int, b: int = BF16, fused: bool = True): """softmax:exp / 减最大值 / 归一化,约 5 次运算/元素。 fused=True 假设整行可驻留片上存储,读 1 遍写 1 遍;不是独立 online normalizer 的通用 IO fused=False 朴素三遍(求 max、求 exp 和、归一化),读 3 遍写 2 遍 """ F = 6.0 * n * d D = b * 2 * n * d if fused else b * (3 * n * d + 2 * n * d) return F, D def attention(N: int, D: int, b: int = BF16, flash: bool = False): """单头注意力。N 是 token 数,D 是每头维度。 flash=False(朴素):分数矩阵 S 和权重 P 都要落 HBM S = Q K^T 写 N^2 P = softmax(S) 读 N^2 写 N^2 O = P V 读 N^2 → 二次项访存 ≈ 4 N^2 b flash=True(强缓存假设下的最低流量,不是真实 FlashAttention IO):Q/K/V 读进来、O 写出去,N^2 项只留在 SRAM → 访存 ≈ 4 N D b """ F = 4.0 * N * N * D + 5.0 * N * N # 两个 N×N×D 的矩阵乘 + softmax if flash: D_bytes = b * (3 * N * D + N * D) # 读 Q K V,写 O else: D_bytes = b * (4 * N * N + 4 * N * D) # 上面再加 S/P 的读写 return F, D_bytes def fmt(x: float, unit: str = "") -> str: for u, s in (("T", 1e12), ("G", 1e9), ("M", 1e6), ("K", 1e3)): if x >= s: return f"{x / s:.2f} {u}{unit}" return f"{x:.2f} {unit}" def main(): print("=" * 78) print("一、四台「机器」的峰值与 ridge point") print("=" * 78) print(f"{'设备':<28s}{'P_peak':>14s}{'beta':>14s}{'I* = P/beta':>16s}") ridges = {} for name, s in SPECS.items(): I_star = s["peak_flops"] / s["peak_bw"] ridges[name] = (s, I_star) print(f"{name:<28s}{fmt(s['peak_flops'], 'FLOP/s'):>16s}" f"{fmt(s['peak_bw'], 'B/s'):>14s}{I_star:>12.1f} FLOP/B") a100 = SPECS["A100-40GB SXM (bf16)"] h100 = SPECS["H100-80GB SXM (bf16)"] print(f"\n A100 -> H100:带宽 ×{h100['peak_bw'] / a100['peak_bw']:.2f}," f"算力 ×{h100['peak_flops'] / a100['peak_flops']:.2f}," f"ridge point {a100['peak_flops'] / a100['peak_bw']:.0f} -> " f"{h100['peak_flops'] / h100['peak_bw']:.0f}(更高 = 更多算子落入带宽受限区)") # ── 二、算子账本 ──────────────────────────────────────────────── N, D = 4096, 128 # 一个典型的 DiT / LLM 单头规模 cases = [ ("逐元素 add [16M]", elementwise(16_000_000, n_read=2)), ("LayerNorm [4096, 3072]", layernorm(4096, 3072)), ("softmax fused [4096, 4096]", softmax(4096, 4096, fused=True)), ("softmax 朴素三遍 [4096, 4096]", softmax(4096, 4096, fused=False)), ("GEMM 512x4096x4096", gemm(512, 4096, 4096)), ("GEMM 4096x4096x4096", gemm(4096, 4096, 4096)), ("attention 朴素 [N=4096,D=128]", attention(N, D, flash=False)), ("attention ideal-min [N=4096,D=128]", attention(N, D, flash=True)), ] print("\n" + "=" * 78) print("二、算子账本(bf16,A100 口径)") print("=" * 78) hdr = (f"{'算子':<32s}{'F (FLOP)':>14s}{'D (Byte)':>14s}" f"{'I':>10s}{'受限':>10s}{'T_roof':>12s}") print(hdr) print("-" * len(hdr)) for name, (F, Db) in cases: I = F / Db T, tc, tm = roofline_time(F, Db, a100["peak_flops"], a100["peak_bw"]) print(f"{name:<32s}{fmt(F):>14s}{fmt(Db, 'B'):>14s}" f"{I:>10.2f}{bound_of(F, Db, a100['peak_flops'], a100['peak_bw']):>10s}" f"{T * 1e3:>10.3f} ms") # ── 三、FlashAttention 把算术强度抬了多少 ──────────────────────── print("\n" + "=" * 78) print("三、物化 S/P 与理想最低 IO 对比(单头 N=4096, D=128, bf16)") print("=" * 78) F_naive, D_naive = attention(N, D, flash=False) F_flash, D_flash = attention(N, D, flash=True) T_naive, _, _ = roofline_time(F_naive, D_naive, a100["peak_flops"], a100["peak_bw"]) T_flash, _, _ = roofline_time(F_flash, D_flash, a100["peak_flops"], a100["peak_bw"]) print(f" FLOPs : {fmt(F_naive)} -> {fmt(F_flash)} " f"(×{F_flash / F_naive:.3f},几乎没变)") print(f" 访存 : {fmt(D_naive, 'B')} -> {fmt(D_flash, 'B')} " f"(×{D_flash / D_naive:.4f},省 {D_naive / D_flash:.1f} 倍)") print(f" 算术强度 : {F_naive / D_naive:.1f} -> {F_flash / D_flash:.1f} FLOP/B " f"(×{(F_flash / D_flash) / (F_naive / D_naive):.1f})") print(f" 受限类型 : {bound_of(F_naive, D_naive, **a100)} -> " f"{bound_of(F_flash, D_flash, **a100)}") print(f" roofline : {T_naive * 1e3:.3f} ms -> {T_flash * 1e3:.3f} ms " f"(×{T_naive / T_flash:.2f})") print(f" 注意:加速来自「少搬 {D_naive / D_flash:.0f} 倍字节」," "不是「少算」——此处仅比较相同主导 FLOPs 与理想最低访存,不是实测内核。") # ── 四、砍算力 vs 砍带宽,谁的收益大 ──────────────────────────── print("\n" + "=" * 78) print("四、同一个 kernel,只翻倍峰值算力或带宽的理想加速") print("=" * 78) print(f"{'算子':<32s}{'I/I*':>10s}{'算力×2 的加速':>16s}{'带宽×2 的加速':>16s}") I_star = a100["peak_flops"] / a100["peak_bw"] for name, (F, Db) in cases: s_flops, s_bw = optimize_headroom(F, Db, a100["peak_flops"], a100["peak_bw"]) print(f"{name:<32s}{(F / Db) / I_star:>10.3f}" f"{s_flops:>14.2f}x{s_bw:>14.2f}x") print("\n 读法:带宽受限的行(I/I* < 1)里「算力×2」那一列几乎都是 1.00——" "\n 在只改变计算峰值且访存不变的理想模型里,该列为 1。") # ── 五、GEMM 什么时候从带宽受限翻到算力受限 ───────────────────── print("\n" + "=" * 78) print("五、GEMM [M,4096] x [4096,4096]:M 多大才开始算力受限") print("=" * 78) K = Nout = 4096 print(f"{'M':>8s}{'I (FLOP/B)':>14s}{'受限':>12s}") prev = None for M in (1, 2, 4, 8, 16, 32, 64, 128, 256, 512, 1024, 4096): F, Db = gemm(M, K, Nout) b = bound_of(F, Db, a100["peak_flops"], a100["peak_bw"]) print(f"{M:>8d}{F / Db:>14.2f}{b:>12s}") if prev == "带宽受限" and b == "算力受限": print(f" ↑ 拐点:M 在 ({M // 2}, {M}] 之间," f"I 越过 ridge point {I_star:.0f}") prev = b print("\n 推论:小 batch 推理(M=1~8)里的 GEMM 是带宽受限的," "这时候做 fp8 量化省的是字节、不是算力。") if __name__ == "__main__": main() memory_ledger.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """显存账本 + batch 扫描:容量这本账怎么算,以及它怎么反过来决定速度。 只依赖 numpy(第二部分要真跑计时)。直接 `python memory_ledger.py` 即可运行,约 10 秒。 第一部分是「容量账」:权重 / KV cache / 激活 / 额外开销四项加总,看 假设 78 GiB 可用预算 能塞多少条并发请求。第二部分是「速度账」:batch 变大之后,同一份权重被更多 token 摊薄,算术强度抬上去,kernel 从带宽受限翻到算力受限——这个拐点用真实 计时测出来,顺便告诉你 batch 该开多大。 """ import json import time from pathlib import Path import numpy as np GIB = 1024 ** 3 def _load_peak(): """优先读 machine_probe.py 存下的标定快照(保证口径一致); 没跑过 machine_probe.py 时退回兜底常数,并照常工作。""" snap = Path(__file__).resolve().parent / "probe_results.json" if snap.exists(): peak = json.loads(snap.read_text())["peak"] return peak["peak_flops"], peak["peak_bw"] return 1680.64e9, 77.21e9 # 兜底:先跑一次 machine_probe.py 更准 P_PEAK, BETA = _load_peak() I_STAR = P_PEAK / BETA def kv_cache_bytes(n_tokens: int, n_layers: int, d_kv: int, b: int = 2) -> int: """KV cache 字节数:每个 token 每层都要存一份 K 和一份 V。 n_tokens 序列长度(或并发请求的总 token 数) n_layers 层数 L d_kv 每层的 KV 总维度 = n_kv_heads × d_head(GQA 时用实际的 KV head 数) b 每个元素的字节数(fp16/bf16 = 2,fp8 = 1) """ return 2 * n_tokens * n_layers * d_kv * b def ledger(n_params: float, b_w: int, n_layers: int, d_kv: int, batch: int, seq: int, act_per_token: int, frag_rate: float): """推理显存账本的四项,单位 GiB。 act_per_token 每个 token 的激活峰值(字节)。推理不保留整层的中间结果, 但当前层的十几份临时张量要同时活着:d_model=4096、fp16 时 一份 [d] 张量 8 KiB,取 16 份 ≈ 128 KiB / token。 这是经验值,随实现(融合程度、是否分块)差距很大。 frag_rate 教学额外开销 / (权重 + KV + 激活),不是实测碎片率。 本函数按未分块 prefill 计 batch*seq 个活跃 token; 逐 token decode 的激活应改为 batch*act_per_token,KV 仍取 batch*seq。 """ w = n_params * b_w kv = kv_cache_bytes(batch * seq, n_layers, d_kv) act = batch * seq * act_per_token frag = (w + kv + act) * frag_rate return dict(weights=w / GIB, kv=kv / GIB, act=act / GIB, frag=frag / GIB, total=(w + kv + act + frag) / GIB) def part1(): print("=" * 80) print("一、显存账本:一张 假设 78 GiB 可用预算,能塞多少条并发请求") print("=" * 80) # Llama-2-7B 的配置:32 层、32 个 KV head、d_head=128 → d_kv = 4096 N_PARAMS, N_LAYERS, D_KV = 7e9, 32, 4096 SEQ = 4096 per_token = kv_cache_bytes(1, N_LAYERS, D_KV) print(f" 模型:7B,fp16 权重 = {N_PARAMS * 2 / GIB:.1f} GiB," f"{N_LAYERS} 层,d_kv = {D_KV}") print(f" KV cache 单价:{per_token / 1024:.0f} KiB / token " f"(一条 {SEQ} 长的序列 = {per_token * SEQ / GIB:.2f} GiB)") print() hdr = (f"{'batch':>6s}{'权重':>9s}{'KV cache':>10s}{'激活':>9s}" f"{'额外开销':>9s}{'合计':>9s}{'78GiB预算':>10s}") print(hdr) print("-" * len(hdr)) for batch in (1, 4, 8, 16, 24, 32): # 朴素预分配:按最大长度全预留,教学假设:以权重+KV+激活总和的 30% 估计额外开销,不是实测碎片比例 naive = ledger(N_PARAMS, 2, N_LAYERS, D_KV, batch, SEQ, act_per_token=131072, frag_rate=0.30) ok = "装得下" if naive["total"] < 78 else "OOM" print(f"{batch:>6d}{naive['weights']:>8.1f}G{naive['kv']:>9.2f}G" f"{naive['act']:>8.2f}G{naive['frag']:>8.2f}G" f"{naive['total']:>8.2f}G{ok:>10s}") # 换成低开销假设(教学假设,非 PagedAttention 实测):额外开销率降到 4% print("\n 同一张卡,把额外开销率从 30% 降到 4%(低开销假设):") print(f"{'batch':>6s}{'朴素合计':>11s}{'低开销假设':>11s}{'多出来的并发':>16s}") for batch in (8, 16, 24, 32): naive = ledger(N_PARAMS, 2, N_LAYERS, D_KV, batch, SEQ, 131072, 0.30) paged = ledger(N_PARAMS, 2, N_LAYERS, D_KV, batch, SEQ, 131072, 0.04) extra = "" if naive["total"] > 78 >= paged["total"]: extra = "从 OOM 变装得下" print(f"{batch:>6d}{naive['total']:>10.2f}G{paged['total']:>10.2f}G" f"{extra:>16s}") print("\n 读法:权重那一项是常数,batch 再大也不变;KV cache 才是随并发") print(" 线性增长的那一项,batch 16 时它已经比权重还大。") print(" 「显存够放权重就能跑」这句话漏掉了后面三项。") # 用 roofline_model 里同款口径:下界 = max(F/P, D/beta)。 # P_PEAK / BETA / I_STAR 已在文件顶部从标定快照读出,这里不再覆盖, # 否则快照刷新后正文引用的数字会和实际跑出来的对不上。 def batch_sweep(repeats: int = 15): """固定一份「权重」,扫 batch,返回每一档的实测延迟/吞吐/算术强度。""" rng = np.random.default_rng(0) d = 4096 W = rng.standard_normal((d, d)).astype(np.float32) # 64 MiB 的「权重」 W_bytes = W.nbytes FP32 = 4 def run(B): x = rng.standard_normal((B, d)).astype(np.float32) y = np.empty((B, d), dtype=np.float32) np.matmul(x, W, out=y) # 预热 best = float("inf") for _ in range(repeats): t = time.perf_counter() np.matmul(x, W, out=y) best = min(best, time.perf_counter() - t) return best rows = [] for B in (1, 2, 4, 8, 16, 32, 64, 128, 256): F = 2.0 * B * d * d D = W_bytes + 2 * B * d * FP32 I = F / D dt = run(B) rows.append(dict(B=B, F=F, D=D, I=I, lat=dt, thr=B / dt, bound="算力" if I >= I_STAR else "带宽", t_pred=max(F / P_PEAK, D / BETA), eff=max(F / P_PEAK, D / BETA) / dt)) return rows, W_bytes def part2(): """batch 扫描实测:延迟、吞吐、算术强度,以及拐点在哪。""" print("\n" + "=" * 80) print("二、batch 扫描实测:吞吐什么时候不再涨") print("=" * 80) rows, W_bytes = batch_sweep() (Path(__file__).resolve().parent / "batch_results.json").write_text( json.dumps(dict(rows=rows, weight_bytes=W_bytes, peak=dict(peak_flops=P_PEAK, peak_bw=BETA)), indent=2)) print(f" 固定「权重」W 形状 [4096, 4096],fp32 = {W_bytes / 2 ** 20:.0f} MiB;" f"batch B 就是一次喂进去的 token 数") print(f" 峰值口径(来自 machine_probe 的标定快照):" f"P = {P_PEAK / 1e9:.0f} GFLOP/s,beta = {BETA / 1e9:.1f} GB/s," f"I* = {I_STAR:.1f}") print() hdr = (f"{'B':>6s}{'F (GFLOP)':>12s}{'D (MB)':>10s}{'I':>9s}{'受限':>8s}" f"{'延迟 ms':>10s}{'吞吐 K/s':>11s}{'下界 ms':>10s}{'达成率':>9s}") print(hdr) print("-" * len(hdr)) prev_bound = None for r in rows: print(f"{r['B']:>6d}{r['F'] / 1e9:>12.3f}{r['D'] / 1e6:>10.2f}" f"{r['I']:>9.2f}{r['bound']:>8s}{r['lat'] * 1e3:>10.3f}" f"{r['thr'] / 1e3:>9.1f}K{r['t_pred'] * 1e3:>10.3f}" f"{r['eff'] * 100:>8.0f}%") if prev_bound == "带宽" and r["bound"] == "算力": print(f" ↑ 拐点:B 从 {r['B'] // 2} 到 {r['B']} 之间" f"越过 ridge point {I_STAR:.1f}") prev_bound = r["bound"] # 反常检测:batch 变大反而变慢,说明库换了代码路径(roofline 看不见这件事) for r0, r1 in zip(rows, rows[1:]): if r1["B"] == 2 * r0["B"] and r1["lat"] > r0["lat"] * 1.5: print(f" ! 反常:B={r0['B']} 只要 {r0['lat'] * 1e3:.2f} ms," f"B={r1['B']} 却要 {r1['lat'] * 1e3:.2f} ms") base = rows[0] top = rows[-1] print(f"\n 延迟:B=1 时 {base['lat'] * 1e3:.3f} ms,B={top['B']} 时 " f"{top['lat'] * 1e3:.3f} ms(涨 {top['lat'] / base['lat']:.1f} 倍)") print(f" 吞吐:B=1 时 {base['thr'] / 1e3:.1f} K/s,B={top['B']} 时 " f"{top['thr'] / 1e3:.1f} K/s(涨 {top['thr'] / base['thr']:.1f} 倍)") print(f" 单样本成本:B={top['B']} 时把 64 MiB 权重的读取摊薄到 {top['B']} 个样本," f"每个样本只摊 {W_bytes / top['B'] / 2 ** 20:.2f} MiB") print(" → 增大 batch 可摊薄权重访问;延迟是否平坦、吞吐能增加多少仍需实测。") print(" → 上面标 ! 的行是 roofline 看不见的东西:同样的公式、更大的 batch," "可能涉及内核选择、调度或计时波动,单凭本表不能确诊。") if __name__ == "__main__": part1() part2() make_figures.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """生成「性能建模与 Profiling」一文的四张解释图。 图里的计时优先读取已保存的实测快照,避免重绘时图文使用不同测量: * roofline.png 峰值与落点来自 machine_probe.calibrate() / probe() * batch_scaling.png 来自 memory_ledger.batch_sweep() * memory_ledger.png 来自 memory_ledger.ledger() * opt_gain.png 来自 roofline_model.optimize_headroom() 只依赖 numpy + matplotlib,直接 `python make_figures.py` 即可运行,约 40 秒。 改了另外三个脚本,这里要重跑,避免图与正文数字对不上。 """ import json import sys from pathlib import Path import matplotlib.pyplot as plt import numpy as np HERE = Path(__file__).resolve().parent sys.path.insert(0, str(HERE)) import machine_probe as MP # noqa: E402 import memory_ledger as ML # noqa: E402 import roofline_model as RM # noqa: E402 OUT = HERE.parent / "figures" try: # exist_ok=True 是必须的:目录已存在时 pathlib 会抛 FileExistsError, # 少数沙箱环境连 exist_ok=True 的 mkdir 也一并拦,这里再兜一层。 OUT.mkdir(exist_ok=True) except PermissionError: if not OUT.is_dir(): raise plt.rcParams.update({ "font.sans-serif": ["PingFang SC", "Arial Unicode MS", "DejaVu Sans"], "axes.unicode_minus": False, "figure.dpi": 160, "savefig.dpi": 160, }) C_MEM = "#e05263" # 带宽受限:红 C_CMP = "#3f7fbf" # 算力受限:蓝 C_ACC = "#2f9e6f" # 好结果:绿 C_GREY = "#94a3b8" # ── 图 1:roofline 曲线 + 实测落点 ──────────────────────────────── def fig_roofline(peak, rows): P = peak["peak_flops"] / 1e9 # GFLOP/s beta = peak["peak_bw"] / 1e9 # GB/s I_star = P / beta # 因为两个都除过 1e9,比值不变 I = np.logspace(-2, 4, 400) perf = np.minimum(P, beta * I) fig, ax = plt.subplots(figsize=(8.4, 5.4)) ax.loglog(I, perf, color="k", lw=2.2) ax.axvline(I_star, color=C_GREY, ls="--", lw=1.4) ax.text(I_star * 1.15, 3, f"ridge point I* = {I_star:.0f}\nFLOP/Byte", color="#475569", fontsize=9, va="bottom") ax.text(0.012, 25, f"峰值算力 P = {P:.0f} GFLOP/s", fontsize=9, color="#475569") ax.text(0.012, 8.5, "斜线 = 带宽天花板,beta = " f"{beta:.0f} GB/s", fontsize=9, color="#475569") # 落点:用实测时间反推实际性能 short = { "逐元素 add(1 遍)": "逐元素 add", "add+relu 两遍(预分配)": "add+relu 两遍", "add+relu 写成一行(有中间数组)": "add+relu 一行写法", "LayerNorm [8192,4096] 预分配": "LayerNorm", "attention 朴素 [N=2048,D=128]": "朴素 attention", "GEMM 512x4096x4096": "GEMM M=512", "GEMM 2048x4096x4096": "GEMM M=2048", "GEMM 4096x4096x4096": "GEMM M=4096", } offsets = { # 手工调过:右边缘和重叠的标签让位 "GEMM M=4096": (-10, 10), "GEMM M=2048": (10, -4), "GEMM M=512": (12, -14), "朴素 attention": (7, -3), } # 左下角四个点挤在一起,用箭头把标签拉到右边的空地上 arrows = { "逐元素 add": (0.42, 9.5), "LayerNorm": (0.42, 5.6), "add+relu 两遍": (0.42, 3.1), "add+relu 一行写法": (0.42, 1.75), } for r in rows: name = short.get(r["name"], r["name"]) y = r["gflops"] ax.scatter([r["I"]], [y], s=52, color=C_MEM if r["bound"] == "带宽" else C_CMP, zorder=5, edgecolor="white", linewidth=0.8) if name in arrows: tx, ty = arrows[name] ax.annotate(name, (r["I"], y), xytext=(tx, ty), fontsize=8.5, color="#334155", arrowprops=dict(arrowstyle="-", color=C_GREY, lw=0.9), va="center") continue dx, dy = offsets.get(name, (7, -3)) ha = "right" if dx < 0 else "left" ax.annotate(name, (r["I"], y), textcoords="offset points", xytext=(dx, dy), fontsize=8.5, color="#334155", ha=ha) ax.set_xlabel("算术强度 I = F / D (FLOP/Byte)") ax.set_ylabel("实测性能 (GFLOP/s)") ax.set_title("图 1:本机的 roofline —— 斜线是带宽,平顶是算力,点是实测落点", fontsize=11) ax.set_xlim(0.01, 5000) ax.set_ylim(0.3, 5000) ax.grid(alpha=0.25, which="both", ls=":") handles = [plt.Line2D([], [], marker="o", ls="", color=C_MEM, label="模型分类:I < I*"), plt.Line2D([], [], marker="o", ls="", color=C_CMP, label="模型分类:I >= I*")] ax.legend(handles=handles, loc="lower right", fontsize=8.5, framealpha=0.9) fig.tight_layout() fig.savefig(OUT / "roofline.png") plt.close(fig) # ── 图 3:batch 扫描的延迟与吞吐 ────────────────────────────────── def fig_batch(rows): B = [r["B"] for r in rows] lat = [r["lat"] * 1e3 for r in rows] thr = [r["thr"] / 1e3 for r in rows] fig, ax1 = plt.subplots(figsize=(8.4, 5.0)) ax1.plot(B, lat, "o-", color=C_CMP, lw=2, label="延迟(左轴)") ax1.set_xscale("log", base=2) ax1.set_xlabel("batch B(一次喂进去的 token 数)") ax1.set_ylabel("延迟 (ms)", color=C_CMP) ax1.tick_params(axis="y", labelcolor=C_CMP) ax2 = ax1.twinx() ax2.plot(B, thr, "s--", color=C_ACC, lw=2, label="吞吐(右轴)") ax2.set_ylabel("吞吐 (K token/s)", color=C_ACC) ax2.tick_params(axis="y", labelcolor=C_ACC) # 拐点:I 越过 ridge point 的地方 cross = next((r for r in rows if r["bound"] == "算力"), None) if cross: ax1.axvline(cross["B"], color=C_GREY, ls=":", lw=1.5) ax1.text(cross["B"] * 1.05, max(lat) * 0.92, f"越过 ridge point\nB ≈ {cross['B']}", fontsize=9, color="#475569") # 库的反常:延迟不随 batch 单调 b8 = next((r["lat"] for r in rows if r["B"] == 8), None) spike = next((r for r in rows if r["B"] in (2, 4) and b8 is not None and r["lat"] > 1.5 * b8), None) if spike: ax1.annotate(f"B={spike['B']} 延迟为 B=8 的 {spike['lat']/b8:.1f} 倍\n(原因需 profile)", (spike["B"], spike["lat"] * 1e3), textcoords="offset points", xytext=(14, -6), fontsize=8.5, color=C_MEM) ax1.set_title("图 3:batch 扫描 —— 实测延迟、吞吐与 roofline 模型交界", fontsize=11) ax1.grid(alpha=0.25, ls=":") lines1, lab1 = ax1.get_legend_handles_labels() lines2, lab2 = ax2.get_legend_handles_labels() ax1.legend(lines1 + lines2, lab1 + lab2, loc="upper left", fontsize=9) fig.tight_layout() fig.savefig(OUT / "batch_scaling.png") plt.close(fig) # ── 图 4:显存账本 ──────────────────────────────────────────────── def fig_ledger(): batches = [1, 4, 8, 16, 24, 32] naive = [ML.ledger(7e9, 2, 32, 4096, b, 4096, 131072, 0.30) for b in batches] paged = [ML.ledger(7e9, 2, 32, 4096, b, 4096, 131072, 0.04) for b in batches] fig, (ax, ax2) = plt.subplots(1, 2, figsize=(11.2, 4.8), gridspec_kw={"width_ratios": [1.35, 1]}) x = np.arange(len(batches)) keys = [("weights", "权重", "#3f7fbf"), ("kv", "KV cache", "#e05263"), ("act", "激活", "#e8a33d"), ("frag", "额外开销", "#94a3b8")] bottom = np.zeros(len(batches)) for k, label, color in keys: vals = np.array([d[k] for d in naive]) ax.bar(x, vals, bottom=bottom, label=label, color=color, width=0.62, edgecolor="white", linewidth=0.6) bottom += vals ax.axhline(78, color="k", ls="--", lw=1.3) ax.text(len(batches) - 0.4, 79, "教学假设:78 GiB 可用预算", fontsize=8.5, ha="right", color="#334155") for i, d in enumerate(naive): if d["total"] > 78: ax.text(i, d["total"] + 2, "OOM", ha="center", fontsize=9, color=C_MEM, fontweight="bold") ax.set_xticks(x) ax.set_xticklabels(batches) ax.set_xlabel("并发序列数 batch") ax.set_ylabel("显存 (GiB)") ax.set_title("图 4a:未分块 prefill 教学账本(额外开销率 30%)", fontsize=10.5) ax.legend(fontsize=8.5, ncol=2) ax.grid(axis="y", alpha=0.25, ls=":") w_naive = [d["total"] for d in naive] w_paged = [d["total"] for d in paged] ax2.bar(x - 0.2, w_naive, width=0.4, label="朴素", color="#e05263") ax2.bar(x + 0.2, w_paged, width=0.4, label="低开销假设", color="#2f9e6f") ax2.axhline(78, color="k", ls="--", lw=1.3) ax2.set_xticks(x) ax2.set_xticklabels(batches) ax2.set_xlabel("并发序列数 batch") ax2.set_ylabel("合计显存 (GiB)") ax2.set_title("图 4b:只把额外开销率从 30% 降到 4%", fontsize=10.5) ax2.legend(fontsize=8.5) ax2.grid(axis="y", alpha=0.25, ls=":") fig.suptitle("图 4:显存账本的四项 —— 权重是常数,KV cache 随并发线性增长", fontsize=11) fig.tight_layout() fig.savefig(OUT / "memory_ledger.png") plt.close(fig) # ── 图 2:砍算力 vs 砍带宽,各有多少钱 ──────────────────────────── def fig_opt_gain(): a100 = RM.SPECS["A100-40GB SXM (bf16)"] ratio = np.logspace(-2, 1.2, 300) sp_flops, sp_bw = [], [] for r in ratio: # 造一个算术强度恰好是 r × I* 的算子:固定 D,F 由 r 决定 D = 1e8 F = r * (a100["peak_flops"] / a100["peak_bw"]) * D s_f, s_b = RM.optimize_headroom(F, D, a100["peak_flops"], a100["peak_bw"]) sp_flops.append(s_f) sp_bw.append(s_b) fig, ax = plt.subplots(figsize=(8.4, 4.8)) ax.semilogx(ratio, sp_flops, color=C_CMP, lw=2.2, label="算力翻 2 倍能拿到的加速") ax.semilogx(ratio, sp_bw, color=C_MEM, lw=2.2, label="带宽翻 2 倍能拿到的加速") ax.axvline(1.0, color=C_GREY, ls="--", lw=1.4) ax.text(1.05, 1.02, "I = I*:分界线", fontsize=9, color="#475569") ax.fill_between(ratio, 0.98, 2.02, where=np.array(ratio) < 1, color=C_MEM, alpha=0.07) ax.fill_between(ratio, 0.98, 2.02, where=np.array(ratio) >= 1, color=C_CMP, alpha=0.07) ax.text(0.05, 1.9, "带宽受限区:\n换更快的算力单元 = 0 收益", fontsize=9, color=C_MEM) ax.text(6, 1.9, "算力受限区:\n加带宽 = 0 收益", fontsize=9, color=C_CMP) ax.set_xlabel("算术强度 / ridge point (I / I*)") ax.set_ylabel("能拿到的加速比") ax.set_ylim(0.95, 2.1) ax.set_title("图 2:投资之前先看这张图 —— 你的 kernel 在分界线哪一侧", fontsize=11) ax.grid(alpha=0.25, ls=":") ax.legend(fontsize=9, loc="center right") fig.tight_layout() fig.savefig(OUT / "opt_gain.png") plt.close(fig) def main(): # --only <图名>:只重画指定一张(roofline / batch / ledger / gain), # 避免为了改一张图的标签把所有实测算子重跑一遍、数字跟正文引用分叉。 only = None if "--only" in sys.argv: only = sys.argv[sys.argv.index("--only") + 1] # 优先读 machine_probe.py 存的标定快照,保证图上数字与正文引用同源 if MP.SNAPSHOT.exists(): snap = json.loads(MP.SNAPSHOT.read_text()) peak = snap["peak"] print(f"读取标定快照({snap['timestamp']},{snap['note']}):") else: print("没找到标定快照,现场标定一遍…") peak = MP.calibrate(verbose=False) MP.SNAPSHOT.write_text(json.dumps( dict(peak=peak, rows=MP.probe(peak, verbose=False), note="fresh", timestamp=""))) print(f" P = {peak['peak_flops'] / 1e9:.0f} GFLOP/s, " f"beta = {peak['peak_bw'] / 1e9:.1f} GB/s, " f"I* = {peak['peak_flops'] / peak['peak_bw']:.1f}") if only in (None, "roofline"): rows = snap.get("rows") if MP.SNAPSHOT.exists() and "snap" in locals() else None if not rows: rows = MP.probe(peak, verbose=False) fig_roofline(peak, rows) if only in (None, "batch"): batch_snap = HERE / "batch_results.json" if batch_snap.exists(): sweeps = json.loads(batch_snap.read_text())["rows"] else: sweeps, weight_bytes = ML.batch_sweep() batch_snap.write_text(json.dumps(dict(rows=sweeps, weight_bytes=weight_bytes, peak=peak), indent=2)) fig_batch(sweeps) if only in (None, "ledger"): fig_ledger() if only in (None, "gain"): fig_opt_gain() if only and only not in ("roofline", "batch", "ledger", "gain"): raise ValueError("unknown figure: " + only) print(f"\n图已写入 {OUT}") if __name__ == "__main__": main()
2026年09月25日
4 阅读
0 评论
0 点赞
2026-09-24
AIGC 基本功|自注意力机制的计算与显存账本-MHA
自注意力机制的计算与显存账本 所属方向:注意力与位置编码 | 难度:入门 | 前置知识:无(这是知识树注意力方向的根节点) 关键词:自注意力、multi-head attention、QKV、复杂度、显存账本、缩放因子 01. 为什么需要它 先看一个真实会撞上的场景:你拿到一个 约 3.62B 参数的简化 Transformer(d_model=3072、32 层),想在 1024×1024 的图生视频任务上做训练,并保留各层激活供反向使用。VAE 八倍下采样、patch size 为 2 之后,单帧 latent 是 128×128、patch 网格是 64×64,单帧序列长度 N=4096(视频若联合多个潜帧,N 还要乘潜帧数)。模型权重 bf16 只有 6.75 GiB,80G 的 A100/H100 看起来绰绰有余,然后第一步就 CUDA OOM。 把账摊开看(怎么算出来的见第 03、04 节):一层注意力在 N=4096 时按下述保守教学账本计为 2.39 GiB 激活,32 层就是 76.5 GiB——是权重的 11 倍。这说明只看权重大小无法判断是否 OOM;在这套训练存储假设下,激活已经超过预算,而且爆的是其中三个特定张量:分数矩阵、softmax 权重、浮点 dropout 乘子,各占 31.4%,三项合计吃掉单层激活的 94%。 按本文简化层结构,N=4096 时两个 N×N 矩阵乘只占整层前向 FLOPs 的 18.2%;它们与线性项之比在 N=18432 达到 1,只看 attention 投影则交叉点为 6144。这是运算量比例,不是运行时间比例;IO、kernel 形状与融合可能让 FLOPs 较少的部分反而更慢。 因此应同时记录 FLOPs、激活存储假设和实际时间线。N=32768 时,本文教学账本的一层存储为 145 GiB,二次项 FLOPs 占 64%;这些值用于理解增长趋势,不能替代实际框架的显存测量。 02. 最小可用理解 三句话讲完核心思想: 每个 token 拿自己的向量生成三份拷贝——Query、Key、Value,然后每个 token 拿自己的 Query 去和所有 token(包括自己)的 Key 做内积打分,分数过 softmax 变成权重,再对所有 Value 加权求和,得到这个 token 的新表示。权重由 Q/K、位置编码和掩码共同决定,输出内容还依赖 V。 计算和显存都分两笔:一笔随 N 线性(QKV 投影、输出投影、FFN),一笔随 N 二次(分数矩阵 S、softmax 权重 P、浮点 dropout 乘子各一份 B×H×N×N)。显存爆炸的几乎都是第二笔,这是显式保存中间量的教学实现假设;融合、重算或不使用 dropout 会改变这笔账。 「O(N²) 是瓶颈」有适用条件:二次项与线性项的比值是 N/(6d),占总量的比例是 N/(6d+N),N 小于 2d 时它连 attention 模块内部的一半都不到。在本例朴素存储方案下,N² 项先成为显存大头;实际耗时瓶颈仍需测量。 这张图要看什么:多头切的是 d_model 这个维度(H·D 拆成 H 份),不是把注意力复制 H 份;切分本身零算术开销,但分数矩阵从 1 个 N×N 变成 H 个 N×N,显存乘上 H。 03. 数学推导 3.1 单头:打分、归一化、加权求和 设输入序列 $X \in \mathbb{R}^{N \times d}$,N 是 token 数,d 是每个 token 的向量维度(d_model)。三个投影矩阵 $W_q, W_k, W_v \in \mathbb{R}^{d \times d}$ 把每个 token 映射成查询、键、值: $$Q = X W_q, \quad K = X W_k, \quad V = X W_v$$ 每个 token 的 Query 要和所有 token 的 Key 算相似度,写成矩阵形式就是一次 $N \times D$ 对 $D \times N$ 的矩阵乘,得到分数矩阵 $S \in \mathbb{R}^{N \times N}$,其中 $S_{ij}$ 是第 i 个 token 对第 j 个 token 的打分: $$S = \frac{Q K^{\top}}{\sqrt{D}}$$ 本小节是单头,H=1、D=d,所以 Q/K/V 都是 N×D;下一小节切成 H 头后才有 D=d/H。除以 $\sqrt{D}$ 不是装饰,推导一下就知道:假设 Q、K 的分量独立、零均值、方差为 1,那么点积的方差是 $$\mathrm{Var}(q \cdot k) = \sum_{i=1}^{D} \mathrm{Var}(q_i k_i) = D$$ 点积的标准差随 $\sqrt{D}$ 线性增长。D=128 时分数的摆动幅度是 D=1 的 11 倍,softmax 拿到这么大的输入会直接饱和:最大的那个分数吃掉几乎全部权重,输出逼近 one-hot。接近 one-hot 时,softmax Jacobian 的多数项会很小;有限 logits 的精确 softmax 通常并非严格 one-hot,浮点舍入可能进一步使梯度消失。除以 $\sqrt{D}$ 恰好把方差归一回 1。第 04 节的代码里有一张 D 从 8 扫到 256 的实测表,饱和是看得见的。 分数过 softmax 变成权重(每行归一化,行内竞争): $$P_{ij} = \frac{\exp(S_{ij})}{\sum_{j'} \exp(S_{ij'})}$$ 最后对 Value 加权求和得到输出 $O = P V$,形状和输入一样是 $N \times d$。工程实现里 softmax 前要先减去每行最大值再取指数,防止 $\exp$ 上溢——这不改变结果,因为分子分母同乘了一个常数。 3.2 多头:切的是维度,不是份数 把 d 维切成 H 段,每段 D = d/H 维当作一个独立的「头」,各算各的注意力,最后拼回来过一个输出投影: $$\mathrm{MHA}(X) = \mathrm{Concat}(\mathrm{head}_1, \dots, \mathrm{head}_H)\, W_o, \qquad \mathrm{head}_h = \mathrm{softmax}\!\left(\frac{Q_h K_h^{\top}}{\sqrt{D}}\right) V_h$$ $Q_h$ 是 $Q$ 的第 h 段 D 列。两个常被搞错的点: 固定 d 时,多头不改变两次矩阵乘的主导 FLOPs。H 个头各做 $N^2 D$ 次乘加,总共 $H \cdot N^2 D = N^2 d$,和不切头(一个 D=d 的单头)一样;softmax、调度等开销仍随头数变化。多头改的是「在多少个独立子空间里同时做注意力」,是表达能力的再分配,不是算力的加倍。 多头增加显存。分数矩阵是按头存的:H 个 $N \times N$。显存里 显式分数存储的 N² 项系数是 H 而不是 1。 3.3 算力账:二次项什么时候过半 一层 transformer 的前向 FLOPs(一次 $[M,K] \times [K,N]$ 矩阵乘算 $2MKN$ 个浮点运算): 四个 d×d 投影(Q、K、V 输入投影 + 输出投影):$8 N d^2$ FFN(升维 4d 再降回):$16 N d^2$ 注意力内部两个 N×N 矩阵乘($QK^{\top}$ 与 $PV$):$4 N^2 d$ 线性项合计 $24 N d^2$,二次项是 $4 N^2 d$,比值等于 $N / (6d)$。令比值等于 1: $$4 N^2 d = 24 N d^2 \quad \Longrightarrow \quad N^{*} = 6d$$ d=3072 时 $N^{*} = 18432$。如果只看 attention 模块内部(4 个投影对 2 个 N×N 乘),交叉点是 $N^{*} = 2d = 6144$。你日常跑的 N=4096 在两条线之下——二次项占整层算力 18.2%,占模块内 40.0%。 3.4 显存账:三个 B×H×N×N 下面采用保守的教学分配模型:同时计入 X/Q/K/V/ctx/O,以及 S、P 和一份浮点 dropout 乘子,全部按 bf16 两字节估算。这不是某个框架的峰值实测,也不是反向传播的最低存储要求;bool mask 通常只占一字节,dropout=0 时可省掉该项,S 通常不必与 P 同时保留。 线性项:X、Q、K、V、加权和、输出,共 6 个 $[B, N, d]$ 张量(不含 FFN 的话); 二次项:分数矩阵 S、softmax 权重 P、浮点 dropout 乘子,各 $[B, H, N, N]$。 $$\text{单层激活} \approx 6 B N d \cdot 2 + 3 B H N^2 \cdot 2 \;\; \text{字节}$$ 两笔相等解得 $N^{*} = 2d/H$,d=3072、H=24 时 *N=256**——序列长度刚过几百,显存就已经被 N² 项主导了。算力和显存的交叉点差 72 倍(18432 对 256),这就是「算力瓶颈来得晚、显存瓶颈来得早」的定量出处。 这张图要看什么:左边 N=4096 的堆叠条里红色三个格子(S、P、浮点 dropout 乘子)占 94%,线性项挤在边上几乎看不见;右边是对数轴,两条教学模型的比值随 N 近似线性增长,N=32768 时为 129 倍。融合侧统一按六份线性张量 6BNd·bytes 计,忽略小的 LSE 与内核工作区;这不是具体框架的峰值比。 04. 代码实现 完整脚本在文末附录(mha_minimal.py、attention_memory.py、flops_ledger.py、make_figures.py),只依赖 numpy。下面按执行顺序拆核心片段,所有数值都是 /usr/local/bin/python3 真跑出来的。 4.1 前向:六行写完 MHA def mha(X, W_q, W_k, W_v, W_o, H, causal=False): B, N, d_model = X.shape D = d_model // H Q = split_heads(X @ W_q, H) # [B, H, N, D] K = split_heads(X @ W_k, H) # [B, H, N, D] V = split_heads(X @ W_v, H) # [B, H, N, D] S = Q @ K.transpose(0, 1, 3, 2) / np.sqrt(D) # [B, H, N, N] if causal: mask = np.triu(np.ones((N, N), dtype=bool), k=1) S = np.where(mask, -np.inf, S) P = softmax(S, axis=-1) # [B, H, N, N] return merge_heads(P @ V) @ W_o, P # [B, N, d_model] 切头函数值得单独看一眼——对这里连续的输入,它是 reshape 加 transpose 的视图变换;非连续输入的 reshape 可能触发拷贝: def split_heads(t, H): B, N, d_model = t.shape D = d_model // H return t.reshape(B, N, H, D).transpose(0, 2, 1, 3) 配置 B=2、N=8、d_model=32、H=4 时各张量的真实形状: [1] 各张量形状 X (2, 8, 32) 输入 Q/K/V (2, 4, 8, 8) 切头后 [B, H, N, D] S (2, 4, 8, 8) 分数矩阵,N×N 是显存爆炸的源头 O (2, 8, 32) 输出,和输入同形 4.2 验证四条性质 脚本跑出来的校验结果: [3] 性质校验 (a) 权重行和为 1 : max|sum(P)-1| = 2.22e-16 (b) 向量化 vs 四重循环 : max|O - O_loop| = 2.22e-16 allclose=True (c) 因果掩码上三角全 0 : max = 0.00e+00 且下三角行和仍为 1 : max|sum-1| = 2.22e-16 (d) 各头权重并不相同 : 头 0 与 H 头平均的平均绝对差 = 0.0223 头之间的平均离散度 : 0.0206(0 表示所有头完全一样) (b) 是最值得做的一次校验:把公式照定义抄成四重循环(一个 token 一个 token 地打分、归一、加权),结果和向量化版本在浮点精度内完全一致。下标搞反、transpose 方向写错这类 bug,靠肉眼很难发现,靠一个慢十倍但显然正确的对照实现能当场抓住。 4.3 缩放因子的实测 D 扫描表(固定 scale=1,看 softmax 最大权重): 固定 scale=1,扫一遍 D 看 softmax 最大权重怎么变: D 分数标准差 最大权重(未缩放) 最大权重(缩放后) 8 2.8587 0.9089 0.2711 16 3.9246 0.9970 0.1718 32 5.6569 1.0000 0.2427 64 8.0203 1.0000 0.1547 128 11.3903 1.0000 0.1322 256 16.0367 1.0000 0.1619 和 3.1 节的推导对上了:分数标准差就是 $\sqrt{D}$(2.8587 ≈ √8,16.0367 ≈ √256)。本次随机样本在 D=16 时最大权重为 0.9970,之后若干行四舍五入显示 1.0000;这说明可能接近饱和,不能证明所有 token 的梯度严格为零。缩放后最大权重回落到 0.13~0.27,分布活着。另外注意主配置(D=8)下的对比:未缩放最大权重 0.5948,缩放后 0.2589——D 小的时候不缩放也能活,缩放用于控制点积随维度增长的方差;这组随机样本不构成某个 D 阈值的通用结论。 4.4 显存账本实跑 attention_memory.py 在 d_model=3072、H=24、B=1、bf16 下逐项清点: [1] 逐项账本 B=1 N=4096 bf16(2 字节) 分数矩阵 S [B, H, N, N] 805,306,368 31.4% softmax 权重 P [B, H, N, N] 805,306,368 31.4% 浮点 dropout 乘子 [B, H, N, N] 805,306,368 31.4% 输入 X / Q / K / V / ctx / O 各 25,165,824 1.0% 合计 2,566,914,048 2.39 GiB → N² 项共 2.25 GiB,线性项共 144.00 MiB,N² 项占 94.1% 随 N 的增长(一层,不含 FFN): [2] N 增长时一层 MHA 的激活显存(B=1, bf16, 含 浮点 dropout 乘子) N 线性项 N² 项 合计 融合侧教学值 倍数 1024 36.00 MiB 144.00 MiB 180.00 MiB 36.00 MiB 5.0x 2048 72.00 MiB 576.00 MiB 648.00 MiB 72.00 MiB 9.0x 4096 144.00 MiB 2.25 GiB 2.39 GiB 144.00 MiB 17.0x 8192 288.00 MiB 9.00 GiB 9.28 GiB 288.00 MiB 33.0x 16384 576.00 MiB 36.00 GiB 36.56 GiB 576.00 MiB 65.0x 32768 1.12 GiB 144.00 GiB 145.12 GiB 1.12 GiB 129.0x 若训练时按同一教学假设保留每层中间量,且不做重算,才可再乘层数:L=32 层时,N=2048 的激活是 20.25 GiB(权重的 3.0 倍),N=8192 是 297 GiB(权重的 44 倍)。推理时不应把当前层临时激活直接乘 L;需要缓存历史的自回归推理另计 KV cache 随 N 和并发数 B 都线性涨:B=32、N=32768、L=80 时光 KV cache 就要 960 GiB,这就是为什么长上下文服务都把 GQA/MLA 当标配。 4.5 算力账本实跑 flops_ledger.py 的占比表(d=3072,含 FFN): N 线性项 二次项 二次项占比 4096 927.71 GFLOPs 206.16 GFLOPs 18.2% 16384 3.71 TFLOPs 3.30 TFLOPs 47.1% 32768 7.42 TFLOPs 13.19 TFLOPs 64.0% 65536 14.84 TFLOPs 52.78 TFLOPs 78.0% 脚本末尾有一段本机 CPU 实测(numpy float32,d_t=1024、16 头,数值每台机器都不同,看趋势): N 投影 (ms) QK^T (ms) 实测比 FLOPs 比 256 1.64 2.36 1.44 0.25 1024 6.64 60.80 9.16 1.00 4096 21.36 1148.15 53.75 4.00 N=256 那行最扎眼:QK^T 的算术量只有投影的四分之一,实测却慢了 1.44 倍。原因在最后一节的算术强度(AI = FLOPs / 访存字节):投影的 AI 是 85~228(权重矩阵被整批 token 反复复用),QK^T 只有 21~31(输出是 H·N²,写完就走)。FLOPs 回答「要做多少运算」,AI 回答「能不能跑快」——这也解释了 FlashAttention 为什么省显存的同时还提速:它压根不把 N² 写回显存,等于把最贵的那笔带宽也省了。 这张图要看什么:两条曲线是二次项算力占比随 N 的爬升,红蓝两条竖虚线分别是 N=2d=6144(只算 attention 模块)和 N=6d=18432(算上 FFN);你常用的 N=4096 在两条线左边很远的位置。 05. 工业级实现对照 参考实现(以 2026-09 的 main 分支为准,上游重构频繁): huggingface/transformers → modeling_llama.py:LlamaAttention.forward huggingface/transformers → masking_utils.py / integrations/flash_attention.py:attention 实现分发 生产代码和第 04 节的最小实现有五处不一样,每一处都有理由。 5.1 不落地 N×N:eager / sdpa / flash 三条路 HF 的 attention 实现 attn_implementation 有三档: eager:与本文显式计算 S/P 的思路相同;最小代码没有 dropout,也不代表训练账本的全部分配。好处是 P 可访问,代价是二次存储;具体峰值取决于 dtype、存活期与 dropout。 sdpa:调 PyTorch 的 scaled_dot_product_attention,由 PyTorch 按硬件、dtype、mask 等选择 flash、memory-efficient 或 math 后端;不能仅凭 sdpa 名称断言没有 N×N 中间量。 flash_attention_2:在线 softmax + 分块计算,显存 O(N·d),本文统一教学账本在 N=32768 时两者为 129 倍,真实节省比例需实测。 训练长序列一律用后两档。代价是 P 不再可见——想可视化注意力图、或给 P 加自定义正则,就得回 eager 或单独导出。 5.2 因果掩码不是加 −inf 的稠密矩阵 最小实现里我建了一个 $N \times N$ 的 bool 矩阵,这本身就又是一笔 N² 显存。生产实现传 is_causal=True 让内核按位置关系现场判断,或者用范围的 sliding_window 参数,掩码矩阵完全不落地。自己手写 causal mask 矩阵是新手常见的第二处 OOM 来源。 5.3 QKV 合并成一个投影 三个 $d \times d$ 投影合并成一个 $d \times 3d$(或直接 qkv_proj),一次 GEMM 出 Q、K、V。算术量不变,但少起两次 kernel、权重读取更连续。代价是 PyTorch 里要自己 chunk(3, dim=-1) 拆回来——本次核对的 modeling_llama.py 仍保留独立的 q_proj/k_proj/v_proj;融合 QKV 是另一些架构或执行后端的选择。 5.4 KV 头数可以比 Q 头少:GQA 多头切分时 K、V 的头数用 $H_{kv} < H$(比如 8 对 32),多个 Query 头共享一组 KV。公式的改动只是把 $K_h$ 换成 $K_{\lfloor h/(H/H_{kv})\rfloor}$。它不改变 attention 核心两次矩阵乘的主导 FLOPs,但可以减少 K/V 投影的 FLOPs,省的是 KV cache 和 KV 的显存与带宽——推理时 KV cache 缩到 $H_{kv}/H$(LLaMA-3 70B 的 8/64 为 1/8)。这是在固定 Query 宽度时减少 KV 头数、压缩线性 KV 项的优化,而 FlashAttention 优化的是 N² 那一笔,两者正交,经常一起用。 5.5 buffer 化与 position_ids 无位置编码且无位置相关掩码的自注意力对 token 排列是等变的。位置信息可通过正弦表、可学习嵌入、RoPE 或掩码注入,其中 RoPE 的旋转点积体现相对位移,不应统称绝对位置编码。生产实现把 cos/sin 表注册成 buffer 预计算缓存,并用外部传入的 position_ids 而不是 arange——因为 KV cache 场景下每个新 token 的位置不是从 0 开始,packed 训练时一段序列内部还要重置。位置怎么进注意力,是 RoPE 那一篇的主题。 06. 代价与边界 把 N² 落地换来了什么,又赔了什么。 朴素实现唯一的优点是 S、P 全程可见:可视化注意力图、或使用依赖完整 P 的蒸馏损失,需要访问相应权重。加到 logits 上的结构化 bias 则不必先物化 P;若内核支持其形式,可在分块时应用。导出完整 P 仍需相应的二次输出存储。工程上常见的折中是:训练用 sdpa/flash,分析时用小 N 的 eager 导出注意力图。 二次项的 FLOPs 收益要按序列长度评估。 若各项运行时间恰好与 FLOPs 成正比,N=4096 时消除占比 18.2% 的二次项,整层理想加速约为 1.22 倍;现实中的 IO 和融合会改变时间占比,因此这不是 FlashAttention 的实测加速上限。 多头的账要两头看。 头数 H 越大,每个头的 D = d/H 越小:每头维度 D 改变会影响子空间容量,但不存在这里能够证明的 D<32 通用质量阈值;固定 d=H·D 时,标准 MHA 的 KV cache 与 KV 投影参数量不因增加 H 而线性增长;若固定 D 则另当别论。所以现代模型反而从「H 越多越好」退到「适度头数 + GQA」:LLaMA-3 70B 用 64 个 Query 头配 8 个 KV 头。 什么时候根本不该用全局自注意力。 像素级 self-attention(把 H×W 个像素当 token)在中等分辨率下 N 就上了万,N² 显存直接不可行——这是 latent diffusion 在压缩空间里计算更经济的原因之一;卷积等其他算子的成本也同时下降。高分辨率密集预测里,窗口注意力、局部注意力是常态而不是妥协。 别只优化注意力。 N 小于 2d 时(d=3072 即 N<6144),attention 模块内部的算力大头是投影;FFN 属于整层的另一部分;显存侧倒是早就归 N² 管。所以「 profiling 之前先改结构」是赌博——performance_profiling 那一篇讲的账本方法就是为此准备的。 07. 经典论文脉络 Attention Is All You Need(Vaswani et al., 2017)——本文锚点。把缩放点积注意力 + 多头定型成今天的形态,丢掉循环结构,端到端只靠注意力。 Neural Machine Translation by Jointly Learning to Align and Translate(Bahdanau et al., 2014)——注意力的史前史:注意力最初是翻译里的一组对齐权重,softmax 那一行的「分布」语义就是从这来的。 Generating Long Sequences with Sparse Transformers(Child et al., 2019)——第一条正面强攻 N² 显存的路线:既然 N² 落不下,就让大部分格子为零。稀疏化一脉的开端。 FlashAttention(Dao et al., 2022)——不动公式、只改执行:在线 softmax + 分块,让 N×N 不落地。精确注意力,不是近似——这一点和稀疏/线性路线有本质区别。 GQA: Training Generalized Multi-Query Transformer Models(Ainslie et al., 2023)——把矛头从 N² 转向 KV cache:KV 头数变少、Q 头分组共享,是许多现代模型降低 KV cache 成本的重要设计。 五篇连起来读的线索:注意力先是「一种对齐手段」(2014),再是「唯一的序列算子」(2017),然后 N² 账单到期,工程上先有人砍格子(2019),再有人改执行不砍精度(2022),最后有人发现真正贵的还有 KV 那笔线性账(2023)。 08. 常见误解 「注意力是 O(N²),所以序列不长时它也是最慢的部分。」 N=4096、d=3072 时二次项只占本文层结构前向 FLOPs 的 18.2%,但不能据此猜测热点。应结合算术强度、硬件与 profiler 判断。 「多头注意力算 H 遍,所以比单头慢 H 倍。」 H·N²·D = N²·d,多头和全维单头的算术量完全相同;多的是 H 份 N×N 显存和 H 份小矩阵乘的调度开销,不是 H 倍算力。 「显存不够就是模型太大,换个更小的模型。」 d=3072、L=32 的模型权重 6.75 GiB,N=8192 时仅激活就 297 GiB。先算激活账(6B·N·d·bytes + 3B·H·N²·bytes 乘层数),再决定动不动模型。这套无重算训练账本提示应先比较开 gradient checkpointing、换 sdpa/flash、或降分辨率。 「除以 √D 是可要可不要的数值技巧。」 本次样本中未缩放 logits 更容易接近饱和;D=16 的最大权重实际为 0.9970,显示 1.0000 也不等于数学上梯度严格为零。它控制点积方差随维度增长;其他归一化或初始化设计也能缓解饱和,不能据此宣称不缩放就一定无法训练。 「浮点 dropout 乘子不占显存。」 bool mask 通常每元素一字节,而 bf16 为两字节;本文账本计的是两字节的浮点 dropout 乘子,两者不能混称。使用逐元素 dropout 的 eager 训练通常还需保留随机掩码或等价信息;占用取决于表示与实现,不一定恰好是一份 bf16 张量。FlashAttention 可以保存随机数状态并在分块反向时重建掩码。 09. 动手验证 把附录里的 mha_minimal.py 存下来直接跑(python mha_minimal.py),对照三处输出: 形状链:X (2,8,32) → Q/K/V (2,4,8,8) → S (2,4,8,8) → O (2,8,32)。确认分数矩阵是按头存的 4 个 8×8,不是 1 个。 缩放对照:未缩放最大权重 0.5948,缩放后 0.2589;再把 D 扫描表看一遍,本次 D=16 的未缩放值是 0.9970,后续多行显示值接近 1。 等价性:向量化实现与四重循环实现的 max 误差应为 1e-16 量级(你机器上具体数字可能略有不同,但 allclose 一定是 True)。 然后做两个改动观察变化: 把 H 从 4 改成 1 再改成 8,跑性质校验 (d):头间离散度会从 0.0206 变成 0(单头没有「头间」可言)——多头不是免费的多样性,是切分带来的。 跑 flops_ledger.py 的 CPU 实测段,找到你机器上「实测比」超过「FLOPs 比」的 N 拐点,和理论交叉点 N*=2d 对一下差多少。 预期最容易翻车的是第二个:很多人会预期实测比从一开始就贴近 FLOPs 比,实际 N=256 时实测比约 1.44、FLOPs 比只有 0.25——算术强度那笔账不在 FLOPs 公式里。 10. 延伸阅读 按知识树的依赖顺序,下一步建议这么走: 旋转位置编码 RoPE 的原理与实现——注意力对位置是盲的,RoPE 是视觉/语言模型目前的主流注入方式;本文的 Q、K 在那里会被旋转一次。 视频 DiT 里的 3D RoPE 与分辨率外推——RoPE 在时间/高/宽三组频率上的拆分,视频生成的位置问题。 性能建模与 Profiling:算力、带宽与显存账本——把本文的两本账变成系统方法,roofline 分析。 序列并行——N² 显存装不下时的横向切分方案。 策略梯度与 PPO 基础 等 RL 方向的文章与本篇无直接依赖,可随时穿插。 附录:完整代码 09 节用到的脚本全文如下(mha_minimal.py、attention_memory.py、flops_ledger.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 mha_minimal.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """多头自注意力(MHA)的最小可运行实现:把公式逐行翻译成 numpy。 只依赖 numpy,直接 `python mha_minimal.py` 即可运行。 刻意不用 PyTorch:讲原理时框架的抽象反而是噪声,而且 numpy 谁都能跑。 生产实现见文章第 05 节对 transformers 的引用。 公式符号与代码变量名的对应关系 ------------------------------ X 输入序列 形状 [B, N, d_model] W_q/W_k/W_v 三个输入投影矩阵 形状 [d_model, d_model] W_o 输出投影矩阵 形状 [d_model, d_model] Q, K, V 查询 / 键 / 值 形状 [B, H, N, D] S 缩放后的注意力分数 形状 [B, H, N, N],S = Q K^T / sqrt(D) P softmax 后的权重 形状 [B, H, N, N] O 注意力输出 形状 [B, N, d_model] 运行后你会看到:每一层的真实 shape、注意力权重的数值范围、 以及四条性质校验(行和为 1 / 与循环实现等价 / 缩放系数的影响 / 因果掩码)。 """ import numpy as np rng = np.random.default_rng(0) def split_heads(t: np.ndarray, H: int) -> np.ndarray: """[B, N, H*D] -> [B, H, N, D]。 多头不是「复制 H 份再算」,而是把 d_model 这个维度切成 H 段, 每段独立算一次注意力,最后再拼回去。这一步只是 reshape + transpose, 本身不做任何算术,但它决定了后面所有矩阵乘的形状。 """ B, N, d_model = t.shape D = d_model // H return t.reshape(B, N, H, D).transpose(0, 2, 1, 3) def merge_heads(t: np.ndarray) -> np.ndarray: """[B, H, N, D] -> [B, N, H*D],split_heads 的逆操作。""" B, H, N, D = t.shape return t.transpose(0, 2, 1, 3).reshape(B, N, H * D) def softmax(x: np.ndarray, axis: int = -1) -> np.ndarray: """数值稳定的 softmax:先减去最大值再取指数,避免 exp 溢出。""" x = x - x.max(axis=axis, keepdims=True) e = np.exp(x) return e / e.sum(axis=axis, keepdims=True) def mha(X: np.ndarray, W_q, W_k, W_v, W_o, H: int, causal: bool = False): """完整的多头自注意力前向。返回 (输出, 注意力权重 P)。""" B, N, d_model = X.shape D = d_model // H # ① 三个投影:每个 token 独立地把自己映射成 query / key / value Q = split_heads(X @ W_q, H) # [B, H, N, D] K = split_heads(X @ W_k, H) # [B, H, N, D] V = split_heads(X @ W_v, H) # [B, H, N, D] # ② 打分:每个 query 和所有 key 做内积,得到 N×N 的分数矩阵 S = Q @ K.transpose(0, 1, 3, 2) / np.sqrt(D) # [B, H, N, N] # ③ 因果掩码:第 i 个 query 只能看 0..i 的 key if causal: mask = np.triu(np.ones((N, N), dtype=bool), k=1) S = np.where(mask, -np.inf, S) # ④ 归一化成权重,再对 V 做加权求和 P = softmax(S, axis=-1) # [B, H, N, N] ctx = P @ V # [B, H, N, D] # ⑤ 拼回 d_model,过输出投影 O = merge_heads(ctx) @ W_o # [B, N, d_model] return O, P def mha_naive_loop(X, W_q, W_k, W_v, W_o, H): """完全不用矩阵乘的「照着定义抄」版本:四重循环。 用来验证上面的向量化实现没有把下标搞反。逻辑等价但慢得多(O(B·H·N²·D))。 """ B, N, d_model = X.shape D = d_model // H Q = (X @ W_q).reshape(B, N, H, D).transpose(0, 2, 1, 3) K = (X @ W_k).reshape(B, N, H, D).transpose(0, 2, 1, 3) V = (X @ W_v).reshape(B, N, H, D).transpose(0, 2, 1, 3) out = np.zeros((B, H, N, D)) for b in range(B): for h in range(H): for i in range(N): # 先逐个算出这一行 N 个分数,再 softmax,再加权求和 s = np.array([float(Q[b, h, i] @ K[b, h, j]) / np.sqrt(D) for j in range(N)]) p = softmax(s) out[b, h, i] = p @ V[b, h] return merge_heads(out) @ W_o def main(): B, N, d_model, H = 2, 8, 32, 4 D = d_model // H print(f"配置: B={B} N={N} d_model={d_model} H={H} D={D}") X = rng.normal(size=(B, N, d_model)) * 0.5 W_q, W_k, W_v, W_o = (rng.normal(size=(d_model, d_model)) / np.sqrt(d_model) for _ in range(4)) # ── 1. 每一层的真实形状 ────────────────────────────────── Q = split_heads(X @ W_q, H) K = split_heads(X @ W_k, H) S = Q @ K.transpose(0, 1, 3, 2) / np.sqrt(D) print("\n[1] 各张量形状") print(f" X {X.shape} 输入") print(f" Q/K/V {Q.shape} 切头后 [B, H, N, D]") print(f" S {S.shape} 分数矩阵,N×N 是显存爆炸的源头") print(f" O {mha(X, W_q, W_k, W_v, W_o, H)[0].shape} 输出,和输入同形") # ── 2. 分数的量级:为什么要除以 sqrt(D) ────────────────── raw = Q @ K.transpose(0, 1, 3, 2) # 不除以 sqrt(D) print("\n[2] 缩放系数 sqrt(D) 的作用") print(f" D = {D}, sqrt(D) = {np.sqrt(D):.4f}") print(f" 未缩放分数的标准差 : {raw.std():.4f} 分布范围 [{raw.min():.2f}, {raw.max():.2f}]") print(f" 缩放后分数的标准差 : {S.std():.4f} 分布范围 [{S.min():.2f}, {S.max():.2f}]") print(f" 未缩放 softmax 的最大权重: {softmax(raw).max():.4f}") print(f" 缩放后 softmax 的最大权重: {softmax(S).max():.4f}") # D 越大,不缩放的分数方差越大,softmax 越容易塌成 one-hot print(" 固定 scale=1,扫一遍 D 看 softmax 最大权重怎么变:") print(" D 分数标准差 最大权重(未缩放) 最大权重(缩放后)") for D_test in (8, 16, 32, 64, 128, 256): q = rng.normal(size=(256, D_test)) k = rng.normal(size=(256, D_test)) s_raw = q @ k.T s_scaled = s_raw / np.sqrt(D_test) print(f" {D_test:4d} {s_raw.std():9.4f} " f"{softmax(s_raw).max():14.4f} {softmax(s_scaled).max():15.4f}") # ── 3. 性质校验 ──────────────────────────────────────── print("\n[3] 性质校验") O, P = mha(X, W_q, W_k, W_v, W_o, H) print(f" (a) 权重行和为 1 : max|sum(P)-1| = " f"{np.abs(P.sum(axis=-1) - 1).max():.2e}") O_loop = mha_naive_loop(X, W_q, W_k, W_v, W_o, H) print(f" (b) 向量化 vs 四重循环 : max|O - O_loop| = " f"{np.abs(O - O_loop).max():.2e} allclose=" f"{np.allclose(O, O_loop)}") _, P_causal = mha(X, W_q, W_k, W_v, W_o, H, causal=True) upper = np.triu(P_causal, k=1) print(f" (c) 因果掩码上三角全 0 : max = {upper.max():.2e}") print(f" 且下三角行和仍为 1 : max|sum-1| = " f"{np.abs(P_causal.sum(axis=-1) - 1).max():.2e}") # 各头学出来的权重并不相同,这才让「多头」有意义 P_mean = P.mean(axis=1) # [B, N, N],把 H 个头的权重平均 diff = np.abs(P[0, 0] - P_mean[0]).mean() spread = np.abs(P[0] - P[0].mean(axis=0, keepdims=True)).mean() print(f" (d) 各头权重并不相同 : 头 0 与 H 头平均的平均绝对差 = {diff:.4f}") print(f" 头之间的平均离散度 : {spread:.4f}(0 表示所有头完全一样)") # ── 4. 不看自己的极端情形:注意力塌缩成「复制」 ────────── print("\n[4] 极端情形:把 K 设成和 Q 完全一样(自注意力且 W_q=W_k)") X2 = rng.normal(size=(1, 4, d_model)) Q2 = split_heads(X2 @ W_q, H) S2 = Q2 @ Q2.transpose(0, 1, 3, 2) / np.sqrt(D) P2 = softmax(S2) diag_share = np.trace(P2[0, 0]) / 4 print(f" 对角线权重占比 = {diag_share:.4f}(1.0 表示每个 token 只关注自己)") if __name__ == "__main__": main() attention_memory.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """自注意力的显存账本:把一层 attention 的每一笔开销都算成字节。 只依赖 numpy(其实只用到它做格式化,核心就是整数乘除)。 直接 `python attention_memory.py` 即可运行。 算的是「训练时一层 MHA 需要为反向传播留下来的激活」, 这是显式保留中间量的保守教学模型,并非框架实测或反向所需最小值。默认的模型尺寸 采用简化 Transformer:d_model=3072、H=24、D=128。 """ import numpy as np BYTES = {"fp32": 4, "fp16": 2, "bf16": 2, "fp8": 1} MIB = 1024 ** 2 GIB = 1024 ** 3 def fmt(nbytes: float) -> str: if nbytes >= GIB: return f"{nbytes / GIB:8.2f} GiB" if nbytes >= MIB: return f"{nbytes / MIB:8.2f} MiB" return f"{nbytes / 1024:8.2f} KiB" def ledger(B: int, N: int, d_model: int, H: int, dtype: str = "bf16", dropout: bool = True): """返回一层 MHA 的逐项激活显存(字节)。 N² 项有三个:分数矩阵 S、softmax 后的 P、以及与所选 dtype 同宽的浮点 dropout 乘子(非 bool mask)。 这三项就是 O(N²) 显存的真身——不是「注意力复杂度是 N²」这句话, 而是本教学模型假设同时保存的三个 B×H×N×N 张量;框架可复用或省略它们。 """ b = BYTES[dtype] linear = B * N * d_model * b # 每个 [B, N, d_model] 的张量 square = B * H * N * N * b # 每个 [B, H, N, N] 的张量 return { "输入 X": linear, "Q": linear, "K": linear, "V": linear, "分数矩阵 S": square, "softmax 权重 P": square, "浮点 dropout 乘子": square if dropout else 0, "加权和 ctx": linear, "输出 O": linear, } def main(): d_model, H, D = 3072, 24, 128 print(f"模型尺寸: d_model={d_model} H={H} D={D} (H*D = {H * D})") # ── 1. 单个 N 下的逐项账本 ────────────────────────────── B, N = 1, 4096 items = ledger(B, N, d_model, H, "bf16") total = sum(items.values()) print(f"\n[1] 逐项账本 B={B} N={N} bf16(2 字节)") print(f" {'项目':<16}{'形状':<22}{'字节数':>14} 占比") for name, nb in items.items(): shape = "[B, N, d_model]" if nb == B * N * d_model * 2 else "[B, H, N, N]" print(f" {name:<16}{shape:<22}{nb:>14,} {nb / total:6.1%}") print(f" {'合计':<16}{'':<22}{total:>14,} {fmt(total)}") quad = sum(v for k, v in items.items() if "S" in k or "P" in k or "乘子" in k) lin = total - quad print(f" → N² 项共 {fmt(quad)},线性项共 {fmt(lin)},N² 项占 {quad / total:.1%}") # ── 2. N 增长时账本怎么变 ─────────────────────────────── print("\n[2] N 增长时一层 MHA 的激活显存(B=1, bf16, 含 浮点 dropout 乘子)") print(f" {'N':>7}{'线性项':>12}{'N² 项':>12}{'合计':>12} " f"{'融合侧教学值':>15} 倍数") for N in (1024, 2048, 4096, 8192, 16384, 32768): it = ledger(1, N, d_model, H, "bf16") q = it["分数矩阵 S"] + it["softmax 权重 P"] + it["浮点 dropout 乘子"] l = sum(it.values()) - q # 与 FlashAttention 文统一:六份线性量的教学预算,忽略小的 LSE 与工作区 flash = 6 * 1 * N * d_model * 2 print(f" {N:>7}{fmt(l):>12}{fmt(q):>12}{fmt(l + q):>12} " f"{fmt(flash):>15} {(l + q) / flash:5.1f}x") # ── 3. 交叉点:从哪个 N 开始 N² 项压过线性项 ──────────── # 3·B·H·N² = 6·B·N·d → N* = 2d / H N_star = 2 * d_model / H print("\n[3] 交叉点") print(f" 3·B·H·N²·bytes = 6·B·N·d·bytes 解得 N* = 2d/H = {N_star:.0f}") print(f" 也就是说 N 超过 {N_star:.0f} 之后,一层 attention 的激活显存") print(f" 就由 N² 项主导;N=4096 时已经超出 {(4096 / N_star):.0f} 倍。") # ── 4. 乘上层数:为什么 L 层比 N 更狠 ──────────────────── print("\n[4] 乘上层数 L(无重算训练,按上述教学模型保留各层中间量)") print(f" {'L':>4}{'N=2048':>12}{'N=4096':>12}{'N=8192':>12} 说明") for L in (12, 24, 32, 48): row = [] for N in (2048, 4096, 8192): tot = sum(ledger(1, N, d_model, H, "bf16").values()) row.append(fmt(tot * L)) print(f" {L:>4}{row[0]:>12}{row[1]:>12}{row[2]:>12}") print(" 以上均未计入模型权重与优化器状态;梯度检查点减少内部保存量,但仍需保留边界激活,并非通用的 1/L。") # ── 5. 推理侧的另一种账本:KV cache ───────────────────── print("\n[5] 推理侧:KV cache 的账本(随 N 线性增长,但随并发数 B 线性增长)") print(" 公式: 2 · B · N · L · H · D · bytes") print(f" {'B':>4}{'N':>7}{'L=32':>12}{'L=80':>12}") for B in (1, 8, 32): for N in (4096, 32768): r = [] for L in (32, 80): nb = 2 * B * N * L * H * D * 2 r.append(fmt(nb)) print(f" {B:>4}{N:>7}{r[0]:>12}{r[1]:>12}") # ── 6. 权重和激活谁更大 ──────────────────────────────── print("\n[6] 权重 vs 激活:别只盯着模型大小") L = 32 # 一层 transformer 的权重:4 个 d×d 投影 + 2 个 FFN 的 4d×d w_per_layer = (4 * d_model * d_model + 2 * d_model * 4 * d_model) * 2 print(f" 一层权重(4 个 d×d + FFN 的 2×4d×d,bf16): {fmt(w_per_layer)}") print(f" L={L} 层总权重 : {fmt(w_per_layer * L)}") for N in (2048, 8192): act = sum(ledger(1, N, d_model, H, "bf16").values()) * L print(f" N={N} 时 L={L} 层激活 : {fmt(act)} " f"(是权重的 {act / (w_per_layer * L):.1f} 倍)") if __name__ == "__main__": main() flops_ledger.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """自注意力的算力账本:把一层 attention 拆成「线性项」和「二次项」两笔。 只依赖 numpy。直接 `python flops_ledger.py` 即可运行。 最后有一段 CPU 实测,耗时约 10 秒;不同机器数值会不同,看趋势即可。 FLOPs 口径:一次矩阵乘 [M,K] @ [K,N] 需要 M·K·N 次乘加(MAC), 一次乘加算 2 个浮点运算,所以 FLOPs = 2·M·K·N。全文统一用这个口径。 """ import time import numpy as np def layer_flops(N: int, d: int, with_ffn: bool = True): """一层 transformer 的前向 FLOPs,拆成线性项和二次项。""" # 4 个 d×d 投影:Q、K、V 三个输入投影 + 1 个输出投影 proj = 4 * 2 * N * d * d # FFN:升到 4d 再降回来,两个矩阵乘 ffn = 2 * 2 * N * d * 4 * d if with_ffn else 0.0 # 注意力内部两个 N×N 的矩阵乘:QK^T 与 P·V attn = 2 * 2 * N * N * d return proj + ffn, attn def fmt_flops(f: float) -> str: for unit, scale in (("P", 1e15), ("T", 1e12), ("G", 1e9), ("M", 1e6)): if f >= scale: return f"{f / scale:7.2f} {unit}FLOPs" return f"{f:7.2f} FLOPs" def main(): d, H = 3072, 24 print(f"模型尺寸: d_model={d} H={H}") # ── 1. 一层里两笔账各是多少 ───────────────────────────── print("\n[1] 一层 transformer 的前向 FLOPs(N=4096, 含 FFN)") linear, quad = layer_flops(4096, d) print(f" 线性项(QKV 投影 + O 投影 + FFN): {fmt_flops(linear)}") print(f" 二次项(QK^T 与 P·V) : {fmt_flops(quad)}") print(f" 二次项占比 : {quad / (linear + quad):.1%}") # 只看 attention 内部:4 个投影 vs 2 个 N×N 矩阵乘 l_attn_only, q_attn_only = layer_flops(4096, d, with_ffn=False) print(f" 只算 attention 模块本身(去掉 FFN): 二次项占 " f"{q_attn_only / (l_attn_only + q_attn_only):.1%}") # ── 2. 交叉点:二次项什么时候超过线性项 ────────────────── # 4N²d = 24Nd² → N* = 6d print("\n[2] 交叉点") print(f" 含 FFN : 4N²d = 24Nd² → N* = 6d = {6 * d}") print(f" 只含投影: 4N²d = 8Nd² → N* = 2d = {2 * d}") print(f" 结论:N 小于 {2 * d} 时,二次项的 FLOPs 少于 attention 投影,") print(f" 这只比较运算量;是否为耗时瓶颈仍取决于访存、内核与硬件。") # ── 3. 占比怎么随 N 变化 ─────────────────────────────── print("\n[3] 二次项占比随 N 变化(d=3072,含 FFN)") print(f" {'N':>7}{'线性项':>16}{'二次项':>16}{'二次项占比':>10}") for N in (512, 1024, 2048, 4096, 8192, 16384, 32768, 65536): l, q = layer_flops(N, d) print(f" {N:>7}{fmt_flops(l):>16}{fmt_flops(q):>16}{q / (l + q):>10.1%}") # ── 4. 训练总算力:前向的三倍 ─────────────────────────── print("\n[4] 训练一个 token 的 FLOPs(前向 + 反向 ≈ 3 倍前向)") for N in (2048, 4096, 8192): l, q = layer_flops(N, d) print(f" N={N:<6} 每 token 前向 {fmt_flops((l + q) / N)}" f" 训练 {fmt_flops(3 * (l + q) / N)}") # ── 5. CPU 实测:二次项的增长是不是真的更快 ────────────── print("\n[5] CPU 实测(本机一次运行,看趋势不看绝对值;numpy float32)") d_t, H_t = 1024, 16 rng = np.random.default_rng(0) X = rng.normal(size=(4096, d_t)).astype(np.float32) W = rng.normal(size=(d_t, d_t)).astype(np.float32) / np.sqrt(d_t) def bench(fn, repeat=3): best = float("inf") for _ in range(repeat): t0 = time.perf_counter() fn() best = min(best, time.perf_counter() - t0) return best * 1000.0 print(f" {'N':>6}{'投影 (ms)':>12}{'QK^T (ms)':>12}{'实测比':>9}" f"{'FLOPs 比':>10}") measured = {} for N in (256, 512, 1024, 2048, 4096): x = X[:N] t_proj = bench(lambda: x @ W) D_t = d_t // H_t Q = (x @ W)[:, :].reshape(N, H_t, D_t).transpose(1, 0, 2) t_qk = bench(lambda: Q @ Q.transpose(0, 2, 1)) # 理论 FLOPs 比:二次项 / 一个投影 ratio_flops = (2 * N * N * d_t) / (2 * N * d_t * d_t) measured[N] = t_qk / t_proj print(f" {N:>6}{t_proj:>12.2f}{t_qk:>12.2f}" f"{t_qk / t_proj:>9.2f}{ratio_flops:>10.2f}") print(f" 注意看 N=256 那一行:QK^T 的算术量只有投影的 0.25 倍,") print(f" 本次 QK^T / 投影耗时比为 {measured[256]:.2f}。耗时比不必等于 FLOPs 比。") # ── 6. 算术强度:FLOPs 回答不了「为什么慢」 ─────────────── print("\n[6] 算术强度 AI = FLOPs / 访存字节(float32,4 字节/元素)") print(" AI 低 = 每读一个字节只做很少的运算 = 带宽先撑不住(访存受限)") print(f" {'算子':<22}{'N=256':>12}{'N=1024':>12}{'N=4096':>12}") rows = [] for N in (256, 1024, 4096): # 投影 X[N,d] @ W[d,d]:读 X 与 W,写输出 f_proj = 2 * N * d_t * d_t m_proj = (N * d_t + d_t * d_t + N * d_t) * 4 # QK^T:读 Q 与 K,写 N×N 的分数矩阵(H 份) f_qk = 2 * N * N * d_t m_qk = (2 * N * d_t + H_t * N * N) * 4 rows.append((N, f_proj / m_proj, f_qk / m_qk)) print(f" {'投影 X@W_q':<22}" + "".join(f"{r[1]:>12.1f}" for r in rows)) print(f" {'注意力 QK^T':<22}" + "".join(f"{r[2]:>12.1f}" for r in rows)) print(f" {'(投影 / QK^T)':<22}" + "".join(f"{r[1] / r[2]:>12.1f}" for r in rows)) print(" 投影的 AI 高出 4~7 倍:权重矩阵被整批 token 复用,读一次能做 N 次乘加。") print(" QK^T 的 AI 低:输出是 H·N²,写完就走,数据复用少。") print(" 这正是 FlashAttention 的第二个收益——不把 N² 写回显存,") print(" 它省的不只是容量,还有带宽。") if __name__ == "__main__": main() make_figures.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """生成「自注意力机制的计算与显存账本」的三张解释图。 数值全部来自同目录下的 attention_memory.py 与 flops_ledger.py, 改了那两个脚本的话这里要跟着重跑,避免图与正文数字不一致。 只依赖 numpy + matplotlib,直接 `python make_figures.py` 即可运行。 """ from pathlib import Path import matplotlib.pyplot as plt import numpy as np from matplotlib.patches import FancyBboxPatch ROOT = Path(__file__).resolve().parents[1] OUT = ROOT / "figures" OUT.mkdir(exist_ok=True) plt.rcParams.update({ "font.sans-serif": ["PingFang SC", "Arial Unicode MS", "DejaVu Sans"], "axes.unicode_minus": False, "figure.dpi": 160, }) C_QUAD = "#e05263" # N² 项:红 C_LIN = "#3f7fbf" # 线性项:蓝 C_FLASH = "#2f9e6f" # FlashAttention:绿 C_GREY = "#94a3b8" INK = "#182238" MIB = 1024 ** 2 GIB = 1024 ** 3 def _style(ax, title, xlabel, ylabel): ax.set_title(title, fontsize=14, weight="bold", color=INK, pad=10) ax.set_xlabel(xlabel, fontsize=11, color="#475569") ax.set_ylabel(ylabel, fontsize=11, color="#475569") ax.tick_params(colors="#475569", labelsize=10) for s in ("top", "right"): ax.spines[s].set_visible(False) for s in ("left", "bottom"): ax.spines[s].set_color("#cbd5e1") ax.grid(axis="y", color="#eef2f7", linewidth=1) ax.set_axisbelow(True) # ──────────────────────────────────────────────────────────── # 图 1:显存账本的构成 # ──────────────────────────────────────────────────────────── def fig_memory_ledger(): d_model, H = 3072, 24 b = 2 # bf16 fig, axes = plt.subplots(1, 2, figsize=(13.2, 4.9), facecolor="#fbfcfe") # 左:N=4096 时的逐项占比(一根堆叠条) ax = axes[0] N = 4096 lin_item = 1 * N * d_model * b quad_item = 1 * H * N * N * b labels = ["输入 X", "Q", "K", "V", "分数矩阵 S", "softmax 权重 P", "浮点 dropout 乘子", "加权和 ctx", "输出 O"] sizes = [lin_item, lin_item, lin_item, lin_item, quad_item, quad_item, quad_item, lin_item, lin_item] colors = [C_LIN] * 4 + [C_QUAD] * 3 + [C_LIN] * 2 total = sum(sizes) left = 0.0 for lab, s, c in zip(labels, sizes, colors): ax.barh([0], [s / GIB], left=left / GIB, color=c, edgecolor="white", linewidth=1.2, height=0.55) if s / total > 0.05: ax.text(left / GIB + s / GIB / 2, 0, f"{lab}\n{s / total:.1%}", ha="center", va="center", fontsize=9.5, color="white", weight="bold") left += s ax.set_xlim(0, total / GIB) ax.set_yticks([]) for s in ("left", "top", "right"): ax.spines[s].set_visible(False) ax.set_xlabel("单层 MHA 的教学激活存储(GiB)", fontsize=11, color="#475569") ax.set_title(f"N={N} 时,三个 N² 项吃掉 94%", fontsize=14, weight="bold", color=INK, pad=10) ax.text(0, -0.42, f"合计 {total / GIB:.2f} GiB | 模型 d_model={d_model}, " f"H={H}, bf16", ha="left", fontsize=10, color="#64748b") # 右:随 N 增长,朴素实现 vs FlashAttention ax = axes[1] Ns = np.array([1024, 2048, 4096, 8192, 16384, 32768]) lin = 6 * Ns * d_model * b # 与逐项账本相同:6 个 [B,N,d] 张量 quad = 3 * H * Ns ** 2 * b # 3 个 [B,H,N,N] 张量 naive = (lin + quad) / GIB flash = (6 * Ns * d_model * b) / GIB # 统一教学线性预算,忽略小的 LSE ax.plot(Ns, naive, "o-", color=C_QUAD, linewidth=2.4, markersize=6, label="朴素实现(落 N² 到显存)") ax.plot(Ns, flash, "s-", color=C_FLASH, linewidth=2.4, markersize=6, label="融合侧教学值(6Nd)") ax.fill_between(Ns, flash, naive, color=C_QUAD, alpha=0.10) ax.set_yscale("log") ax.set_xscale("log") ax.set_xticks(Ns) ax.set_xticklabels([f"{n:,}" for n in Ns], fontsize=9) ax.legend(fontsize=10, frameon=False, loc="upper left") _style(ax, "两者比值随 N 近似线性增长", "序列长度 N(token 数)", "单层教学激活存储(GiB,对数轴)") ax.text(0.98, 0.06, f"N=32768 时相差 {naive[-1] / flash[-1]:.0f} 倍", transform=ax.transAxes, ha="right", fontsize=10, color=C_QUAD, weight="bold") fig.suptitle("自注意力的显存账本:钱花在哪一笔", fontsize=17, weight="bold", color="#0f172a", y=1.03) fig.tight_layout() fig.savefig(OUT / "memory_ledger.png", bbox_inches="tight", facecolor=fig.get_facecolor()) plt.close(fig) # ──────────────────────────────────────────────────────────── # 图 2:算力账本里二次项的占比 # ──────────────────────────────────────────────────────────── def fig_flops_ratio(): d = 3072 Ns = np.logspace(9, 16.5, 300, base=2) # 512 ~ 92682 quad = 4 * Ns ** 2 * d lin_proj = 8 * Ns * d * d # 只含 4 个投影 lin_all = lin_proj + 16 * Ns * d * d # 再算上 FFN fig, ax = plt.subplots(figsize=(8.6, 5.0), facecolor="#fbfcfe") r_only = quad / (quad + lin_proj) r_all = quad / (quad + lin_all) ax.plot(Ns, r_only, "-", color=C_QUAD, linewidth=2.6, label="只算 attention 模块(4 个投影 vs 2 个 N×N 乘)") ax.plot(Ns, r_all, "-", color=C_LIN, linewidth=2.6, label="算上 FFN(整层的线性项)") ax.axhline(0.5, color=C_GREY, linestyle="--", linewidth=1.2) ax.text(Ns[0], 0.52, "50% 线", fontsize=10, color=C_GREY) for N_star, color, tag in ((2 * d, C_QUAD, "N* = 2d = 6,144"), (6 * d, C_LIN, "N* = 6d = 18,432")): ax.axvline(N_star, color=color, linestyle=":", linewidth=1.6) ax.text(N_star * 1.05, 0.06, tag, fontsize=10.5, color=color, weight="bold", rotation=90, va="bottom") ax.set_xscale("log", base=2) ax.set_xticks([512, 1024, 2048, 4096, 8192, 16384, 32768, 65536]) ax.set_xticklabels(["512", "1K", "2K", "4K", "8K", "16K", "32K", "64K"]) ax.set_ylim(0, 1) ax.legend(fontsize=10.5, frameon=False, loc="upper left") _style(ax, "二次项的 FLOPs 占比随 N 增长(不等于耗时占比)", "序列长度 N(token 数)", "二次项占该层前向 FLOPs 的比例") ax.text(0.98, 0.30, "常用的 N=4096:\n只算 attention 时二次项占 40%,\n算上 FFN 后只占 18%", transform=ax.transAxes, ha="right", fontsize=10, color="#475569", bbox=dict(boxstyle="round,pad=0.4", facecolor="#f1f5f9", edgecolor="#e2e8f0")) fig.suptitle("算力账本:二次项什么时候才是主角", fontsize=16, weight="bold", color="#0f172a", y=1.00) fig.tight_layout() fig.savefig(OUT / "flops_ratio.png", bbox_inches="tight", facecolor=fig.get_facecolor()) plt.close(fig) # ──────────────────────────────────────────────────────────── # 图 3:多头到底切了什么 # ──────────────────────────────────────────────────────────── def fig_head_split(): fig, ax = plt.subplots(figsize=(11.4, 5.0), facecolor="#fbfcfe") ax.set_xlim(0, 1) ax.set_ylim(0, 1) ax.axis("off") def block(x, y, w, h, title, lines, color, fs=11): ax.add_patch(FancyBboxPatch( (x, y), w, h, boxstyle="round,pad=0.01,rounding_size=0.02", facecolor=color, edgecolor="white", linewidth=1.6)) ax.text(x + w / 2, y + h * 0.74, title, ha="center", va="center", fontsize=fs + 2, weight="bold", color=INK) ax.text(x + w / 2, y + h * 0.33, "\n".join(lines), ha="center", va="center", fontsize=fs, color="#334155", linespacing=1.5) def arrow(x1, y, x2, label): ax.annotate("", xy=(x2, y), xytext=(x1, y), arrowprops=dict(arrowstyle="-|>", color=C_GREY, linewidth=2.0, mutation_scale=16)) ax.text((x1 + x2) / 2, y + 0.045, label, ha="center", fontsize=11, color="#475569", weight="bold") # X block(0.02, 0.30, 0.16, 0.42, "X", ["[B, N, d_model]", "d_model = H · D"], "#dbeafe") arrow(0.185, 0.51, 0.245, "三个投影") # Q/K/V 未切头 block(0.25, 0.30, 0.17, 0.42, "Q, K, V", ["各 [B, N, H·D]", "仍是完整宽度"], "#e0e7ff") arrow(0.425, 0.51, 0.485, "切头") # 切头后 block(0.49, 0.30, 0.17, 0.42, "切头之后", ["[B, H, N, D]", "reshape + transpose"], "#ede9fe") arrow(0.665, 0.51, 0.725, r"$QK^\top/\sqrt{D}$") # 分数矩阵 block(0.73, 0.22, 0.25, 0.58, "分数矩阵 S", ["[B, H, N, N]", "H 个 N×N,不是 1 个", "← 显存就花在这里"], "#fecaca", fs=11) # 底部说明 ax.text(0.5, 0.10, "每个头只用自己的 D = d_model / H 维去做内积,H 个头各算各的 N×N;\n" "最后把 H 份 D 维结果拼回 d_model,再过一次输出投影 W_o。", ha="center", va="center", fontsize=11.5, color="#475569", linespacing=1.6) ax.text(0.5, 0.02, "关键:分头不改变总算术量(H · N² · D = N² · d_model),它改变的是每个子空间的表达能力。", ha="center", va="center", fontsize=11, color=C_QUAD, weight="bold") fig.suptitle("多头注意力切的是 d_model 这一维,不是复制 H 份", fontsize=16, weight="bold", color="#0f172a", y=0.99) fig.tight_layout() fig.savefig(OUT / "head_split.png", bbox_inches="tight", facecolor=fig.get_facecolor()) plt.close(fig) if __name__ == "__main__": fig_memory_ledger() fig_flops_ratio() fig_head_split() print(f"已生成 3 张图到 {OUT}") for p in sorted(OUT.glob("*.png")): print(f" {p.name} {p.stat().st_size / 1024:.0f} KiB")
2026年09月24日
4 阅读
0 评论
0 点赞
2026-09-21
AIGC 基本功|音频质量评测:MOS、PESQ 与 FAD 各测什么-AudioEval
音频质量评测:MOS、PESQ 与 FAD 各测什么 所属方向:评测与指标 | 难度:工程实战 | 前置知识:无 关键词:音频质量、音频评测、MOS、PESQ、STOI、SI-SDR、FAD、Fréchet Audio Distance、感知评测、语音合成、音频codec、无参考评测 01. 为什么需要它 先看一个教学假设:团队比较两个音频 codec 版本,一个版本对齐更准,SI-SDR 更高,却同时新增了低能量的高频纯音。波形误差下降和听感变差可以同时发生。第 04 节用可复现合成信号展示这个指标盲区;本文没有受试者听测,因此不把「刺耳」或「无差别」写成测得的事实。7~7.5 kHz 的窄带成分可能可闻,其显著程度取决于响度、掩蔽、设备和听者;不能把这一段称为人耳最敏感频段。 这个问题不是「评测做少了」,而是指标回答的问题和团队想问的问题不是同一个。SI-SDR 回答的是「输出波形和参考波形逐采样点对得上吗」,团队想问的是「用户听着难受吗」。这两个问题之间的距离,就是这篇要讲的东西。 更麻烦的是,音频指标的「鄙视链」比图像指标更长、更碎:PESQ 面向语音质量,含窄带与宽带版本;STOI 面向语音可懂度,音乐生成用的是分布级的 FAD,而人工 MOS 需要组织听测,并报告受试者、协议与置信区间。没有一个指标是免费的午餐——先搞清楚每个指标在测什么,再决定用哪个,比记住任何一行数字都重要。 02. 最小可用理解 先把这篇要用的缩写认全——音频圈的缩写比图像圈更碎,不先对齐名词,后面每个数字都容易理解错: 缩写 全称 中文 一句话 MOS Mean Opinion Score 平均意见分 找一群人听、打分取平均。按指定听测协议汇总主观评分,有成本和统计不确定性 PESQ Perceptual Evaluation of Speech Quality 语音质量感知评估 语音有参考质量指标,P.862 为窄带、P.862.2 为宽带扩展 STOI Short-Time Objective Intelligibility 短时客观可懂度 只测「内容听不听得清」,不测好不好听 SI-SDR Scale-Invariant Signal-to-Distortion Ratio 尺度不变信号失真比 允许整体缩放后再比波形,衡量逐采样点对齐程度 FAD Fréchet Audio Distance Fréchet 音频距离 不需逐条配对参考,仍需参考音频集合,比较两堆特征分布,思想来自图像的 FID STFT Short-Time Fourier Transform 短时傅里叶变换 把波形切帧做 FFT,得到随时间变化的频谱 SSL Self-Supervised Learning 自监督学习 UTMOS 这类 MOS 预测器用来提特征的预训练模型 DNSMOS / UTMOS — — 用模型预测 MOS 的两种常见方案(前者微软、后者东京大学 SaruLab) 三句话讲完框架: 每个指标只回答一个具体的问题:「波形逐点对齐了吗」(SI-SDR/SNR)、「频谱形状像吗」(对数梅尔距离)、「这两堆音频听起来像同一个世界来的吗」(FAD)、「人听着舒服吗」(MOS,以及预测它的 UTMOS/DNSMOS)。问题不同,答案就不同——不存在「最好的指标」,只存在「问对了问题的指标」。 有参考和无参考是两大阵营。有参考指标(SI-SDR、PESQ、梅尔距离)需要对应的参考信号,适合 codec、增强、分离等任务;逐条无参考指标(UTMOS、DNSMOS)不需要对应参考;FAD 不需要逐条配对,但需要代表目标分布的参考集合。是否适合音乐或音效仍取决于训练域和特征器。 客观指标是 人耳的代理,代理就有失真。失真最大的地方往往是「在部分听测条件下不显著、指标却敏感」的维度(波形相位、时移、整体音量),和「人耳一耳朵就听出来、指标却毫无反应」的维度(高频毛刺、音色瑕疵)。第 04 节会用一张表把这两类失真都造出来给你看。 03. 数学推导 3.1 SI-SDR:为什么它对音量不敏感、对时移却如此脆弱 信号处理里的「失真」最朴素的定义是信噪比: $$\mathrm{SNR} = 10 \log_{10} \frac{\| s \|^{2}}{\| \hat{s} - s \|^{2}}$$ 其中 $s$ 是参考信号、$\hat{s}$ 是待测信号,范数就是逐采样点的能量。它有一个明显的不合理之处:如果输出和参考只差一个整体音量($\hat{s} = 0.95\,s$),按理说听感几乎没差别,但 SNR 会把它记成 $20\log_{10}(1/0.05) \approx 26$ dB 的「失真」。 SI-SDR 的解法是先做最优缩放,再算残差——这也是它名字里 Scale-Invariant 的来历: $$\alpha = \frac{\langle \hat{s}, s \rangle}{\langle s, s \rangle}$$ $$\mathrm{SI\text{-}SDR} = 10 \log_{10} \frac{\| \alpha s \|^{2}}{\| \hat{s} - \alpha s \|^{2}}$$ $\alpha$ 是把待测信号投影到参考方向上的最小二乘系数:分子是内积 $\langle \hat{s}, s \rangle$(两个波形逐点相乘再求和,衡量「同相程度」),分母是参考信号的能量,用来归一化。缩放完再算残差,整体音量的差异就被精确抵消了。 但这个设计同时埋了雷。SI-SDR 衡量的是逐采样点的对齐程度,而波形对齐对时间平移极端敏感:把信号往后挪 8 个采样点(16 kHz 采样率下只有 0.5 毫秒,单独播放时通常不易觉察),波形的峰谷就全部错开了,内积 $\langle \hat{s}, s \rangle$ 急剧下降、残差急剧上升,SI-SDR 可以从正几十分贝直接掉到负值——第 04 节的实验里,0.5 毫秒的时移让 SI-SDR 掉到了 $-6.25$ dB。 所以使用 SI-SDR 前要确认与参考的对齐约定(降噪、回声消除、codec 重建),并且要明白它测的是去掉全局增益后的波形误差,对齐只是影响因素之一。 顺带回答一个高频疑问:为什么叫 SDR 不叫 SNR? 因为这个指标不是从「信噪比」这棵树上长出来的,而是从盲源分离那套评价体系里长出来的。2006 年的 BSS_eval 工具箱(Vincent / Gribonval / Févotte,Performance measurement in blind audio source separation)把估计信号相对参考的总误差拆成四类: 残留干扰(interference):别的声源串进来多少 加性噪声(noise):环境噪声剩多少 算法伪影(artifacts):算法自己制造出来的人工痕迹 通道误差(channel errors):增益、滤波这类整体偏差 对应地给出 SIR / SAR / ISR 三个分项,以及一个总指标——SDR(Signal-to-Distortion Ratio)。 关键就在 Distortion 这个词:它涵盖上面四类,而 SNR 的「噪声」只对应其中一类。在盲源分离里,最伤听感的往往不是噪声,而是别的源串进来的干扰和算法自己的伪影,叫「信噪比」会名不副实。 后来 2019 年 Le Roux 等人在 SDR – half-baked or well done?(arXiv:1811.02508)里指出 BSS_eval 版 SDR 有两个致命问题:bss_eval_sources 允许用 512 阶滤波器去「修改参考信号」来拟合估计值,于是把某些频段直接置零也能拿到近乎无穷的 SDR;bss_eval_images 的 SDR 又退化成了普通 SNR,不允许全局缩放,反被一些算法无意间靠调音量刷了分。他们给出的修法就是本文用的 SI-SDR——误差项与参考信号严格正交,彻底去掉对幅度缩放的依赖,命名沿袭 SDR 保持与社区的连续性。 SDR 的 distortion 可以包括噪声、串扰和伪影等误差。SNR 与 SI-SDR 是否数值相等还取决于缩放、投影和预处理约定,不能仅凭「只有加性噪声」判定。工程上 SI-SNR 常用于同类投影指标,比较实现时须检查是否先去均值、如何处理静音。 3.2 对数梅尔距离:把「人耳怎么听」写进公式 梅尔刻度近似描述感知音高与频率的非线性关系,不能据此断言某两个高频纯音「人耳几乎分不清」。本文采用 HTK 风格的映射: $$m(f) = 2595 \log_{10}\left(1 + \frac{f}{700}\right)$$ 其中 $f$ 的单位为 Hz。短时傅里叶变换(STFT)得到幅度谱后,用三角滤波器聚合到梅尔频带,再取对数。本文的最小实现是梅尔幅度谱,不是常见的梅尔功率谱;两者必须区分。令 $M(k,t)$ 为第 $k$ 带、第 $t$ 帧的幅度,以每条音频的谱峰值 $M_{\max}$ 归一化,定义: $$L(k,t)=\ln\max\left(M(k,t)/M_{\max},10^{-4}\right)$$ $$D_{\mathrm{mel}} = \frac{1}{KT}\sum_{k,t}\left|L(k,t)-\hat L(k,t)\right|$$ $K,T$ 分别为梅尔带数和帧数。这里的单位是自然对数幅度差,不是 dB;换成 dB 要乘 $20/\ln 10$。独立归一化是为了在本实验里忽略整体增益,若任务关心响度,就不能先把这类差异消掉。幅度下限 $10^{-4}$ 对应峰值下 80 dB,防止对接近零的值取对数;若使用功率谱,同样的 80 dB 下限应是 $10^{-8}$。窗口、hop、滤波器、幅度/功率约定、floor 和归一化都会改变数值,需完整报告。 3.3 Fréchet Audio Distance:从「两条音频」到「两堆音频」 前面两个指标都需要一条参考音频。但音乐生成、音效生成没有「正确答案」——你没法说这段生成的钢琴曲应该长什么样。这类任务需要无参考指标,FAD 是其中最常用的一个。 FAD 的思路借自图像领域的 FID:不比较单条音频,而是比较两堆音频的分布。把每条音频过一个特征提取器(原论文用 VGGish,得到 128 维特征),每堆音频的特征就变成了高维空间里的一团点云,然后近似成高斯分布,问:这两个高斯离得多远? FAD 使用两个拟合高斯之间的 2-Wasserstein 距离平方(沿用 Fréchet distance 名称): $$\mathrm{FD} = \left\| \mu_{r} - \mu_{t} \right\|^{2} + \mathrm{Tr}\left( \Sigma_{r} + \Sigma_{t} - 2 \left( \Sigma_{r} \Sigma_{t} \right)^{1/2} \right)$$ $\mu_{r}, \mu_{t}$ 是参考集和测试集的特征均值(每维上特征的平均),$\Sigma_{r}, \Sigma_{t}$ 是协方差矩阵(特征之间的联合波动结构),$\mathrm{Tr}$ 是矩阵的迹(对角线元素之和)。直觉上:均值差衡量「两团云的中心差多远」,协方差项衡量「两团云的形状差多远」。 样本量会影响均值和协方差的估计。当样本数不超过特征维度时,中心化样本协方差必然秩亏;但秩亏不表示距离无法计算。更关键的是估计偏差和方差,其大小取决于特征谱、独立样本数和参考集,而非单由维度决定。第 04 节的 16 维合成例子中,同分布两组的距离从 8 个样本时的 13.1390 降到 300 个样本时的 0.1844;这是一组随机种子的结果,不是通用阈值。 实战里怎么检查样本量的影响 报告同分布基线及其波动。按原始音频为单位反复随机分组,让组大小与真正的比较一致,观察距离分布。不能拿一次基线的「2 倍」当显著性阈值;应结合效应量、重采样区间或预先设计的置换检验。 不要把切片数当独立样本数。0.96 秒窗口是 VGGish 常见的输入单位。帧级统计与 clip 级统计评价的对象不同,同一条音频的窗口又高度相关。切片可以提供局部信息,但不会凭空产生同等数量的独立录音;bootstrap 应按音频或更高层级成组抽样,不能把相关帧当成独立样本。 固定评测协议。报告参考集、生成集规模、总时长、窗口步长、特征器和归一化。更换这些条件后,旧 FAD 数字不再是同一把尺子。 原始 FAD 论文 与 后续 FAD 评测研究 讨论了特征器和样本量偏差。本次修订移除了缺少原始音频与评测脚本的历史「真实案例」图表;下面保留可由附录完整复现的合成实验,避免把不可复核的数值当作结论依据。 04. 代码实现 完整脚本在文末附录,这里只看核心片段与真实输出。数值脚本只用 numpy;配图脚本另需 matplotlib,不依赖音频模型权重。 第一段:构造三种不同机制的退化(metric_disagreement.py)。干净信号用「基频 + 6 个谐波 + 慢包络」合成;三种退化分别是音量 ×0.95、叠 7/7.5 kHz 纯音、循环时移 0.5 ms: def degrade(x, kind, sr=SR): if kind == "A": return 0.95 * x # 音量差 5% if kind == "B": # 7/7.5 kHz 纯音,均低于 16 kHz 采样率的 Nyquist 频率 t = np.arange(len(x)) / sr buzz = 0.02 * ( np.sin(2 * np.pi * 7000 * t) + np.sin(2 * np.pi * 7500 * t) ) return x + buzz if kind == "C": return np.roll(x, 8) # 时移 8 采样 = 0.5ms raise ValueError(kind) 第二段:SI-SDR 的最小实现(缩放抵消就三行): def si_sdr(ref, est): """Scale-Invariant SDR。 先把 est 投影到 ref 上(找一个最优缩放 alpha),剩下的残差才算失真。 所以「整体音量变了」在 SI-SDR 里被完全抵消——这是设计意图, 但也意味着它测不出增益错误。 """ if np.dot(ref, ref) <= 0 or np.dot(est, est) <= 0: raise ValueError("SI-SDR needs non-silent reference and estimate") alpha = np.dot(est, ref) / np.dot(ref, ref) e_target = alpha * ref e_res = est - e_target num = np.dot(e_target, e_target) den = np.dot(e_res, e_res) if den <= 1e-24 * num: return float("inf"), alpha return 10 * np.log10(num / den), alpha 真实运行输出(16 kHz、1 秒合成乐音): 操作 SI-SDR(dB) SNR(dB) 梅尔距离(ln幅度) 高频占比 高频谱平坦度 A 音量×0.95 inf 26.02 0.0000 3.808e-14 0.3256 B 7/7.5kHz纯音 26.58 26.58 0.3472 2.188e-03 近0 C 循环时移0.5ms -6.25 -0.51 0.0187 3.836e-14 0.3272 A 的全局增益在 SI-SDR 和独立归一化的梅尔谱里被消除;SNR 仍会计入它。$\alpha=0.95$ 表示待测信号为参考的 0.95 倍,并不是把 0.95「倒过来」。 B 的 SI-SDR 为 26.58 dB,但它不是主观音质的合格线。高频能量占比明显上升,说明多了高频成分;这些成分是否不悦耳,需要听测。 C 在波形上错位,SI-SDR 大幅下降,梅尔距离较小。这里 np.roll 是循环时移,会把尾部搬到开头;它只是对齐敏感性的演示,真实系统延迟应在共同有效区间评测。 谱平坦度描述所选频带的谱能量是否集中。纯音通常低、平坦噪声通常高,但音乐本来就含纯音,白噪声也可能是缺陷,因此不能用统一的 0.05 阈值判好坏。本实验干净信号的高频能量约为 $10^{-14}$ 量级,已几乎为空;此时平坦度主要由窗函数泄漏和数值 floor 决定,0.3256 不是「健康音频标准」。高频能量占比与平坦度在理想非零计算下都对全局增益不变,实际小能量频带则应先做绝对能量门控。 第三段:FAD 与样本量(fad_sample_size.py)。没有 VGGish,就用 16 维归一化对数梅尔幅度均值当特征——Fréchet 距离这一步的数学完全一样。核心实现采用对称半正定夹心矩阵。不能把一般非对称的 $\Sigma_r\Sigma_t$ 直接交给只处理对称矩阵的 eigh;通用矩阵平方根算法可以处理乘积,但不是这里的最小实现。附录检查非半正定输入,并仅把容差内的负舍入误差截为零。 真实运行输出(每行重新抽取样本,固定随机种子): 每堆样本数 同分布距离 带纯音距离 8 13.1390 47.5224 16 3.8499 49.1996 32 1.6452 46.6619 64 0.9992 46.3322 128 1.1044 45.4964 300 0.1844 45.4057 总体分布相同才有理论距离零,两个有限样本的估计值不必为零,也不必随样本量严格单调下降。这里是自定义 16 维梅尔特征的高斯距离演示,不是标准 VGGish FAD。 配图由 mel_spectrogram.py 重绘,单位为 dB,和表中自然对数单位相差固定倍数。B 的两条纯音均低于 Nyquist 频率;旧实现的 sin(2π·8000t) 在 16 kHz 采样时落在零点,不能用它模拟一个存在的 8 kHz 纯音。 图要看什么:B 在高频多出亮带;A 因独立归一化与参考重合;C 的谱形状变化较小。色谱只能证实频谱结构变化,不能代替主观听测。 05. 工业级实现对照 最小实现是为了讲清原理,生产里的实现多出的主要是「工程鲁棒性」: FAD:TorchEval 的 FrechetAudioDistance 的 FrechetAudioDistance(以 2026-09 的实现为准)把特征提取抽象成 preproc + model 两个可插拔组件,with_vggish() 一行构建官方配置;Microsoft 的 fadtk 则把 CLAP、MERT、HuBERT、Whisper 等十来种特征提取器都接了进来——FAD 的数值和特征提取器强绑定,换提取器等于换了一把尺子,不同论文的 FAD 数字不能直接比。另外它用流式更新(update/compute 两段式),避免一次把全部音频载入内存。 编解码器:facebookresearch/encodec/encodec/model.py 的 EncodecModel 提供 encode/decode 与可调目标码率的残差量化器,评测 codec 时要用它官方的 24 kHz / 48 kHz 配置,因为不同采样率下量化器的带宽分配不同,失真谱型也不同。 PESQ / STOI:两者来源不同。PESQ 对应 ITU-T P.862 系列,P.862.2 是宽带扩展;ITU 已于 2024-01-05 撤回该系列,并指向 P.863 系列。历史 benchmark 仍可报告 PESQ,但要写清版本、带宽和采样率。STOI 则来自 Taal 等人的可懂度研究,不是 P.862 标准;作者资料与代码 给出来源。常用实现为 pesq / pystoi,使用时仍须核对模式与预处理。 MOS 预测:UTMOS/DNSMOS 可以对单条样本输出预测分,适合筛选和回归测试。域内和域外的绝对值、排序都应在目标数据上验证,不能保证域内排序必然正确。人工听测也并非不可复现,可按 ITU-T P.800 等协议固定设计并报告评分分布和置信区间。 06. 代价与边界 指标 适合的问题 什么时候会骗你 SI-SDR 时间严格对齐的重建任务 对时移和相位差敏感;对单个全局增益不敏感 SNR 同上,且关心音量 对感知无感的相位问题同样过敏 对数梅尔距离 音色/频谱保真度 归一化会消除整体增益;逐帧比较仍对时序变化敏感,时间聚合又会丢失结构 PESQ / STOI 语音质量 / 语音可懂度 需按版本匹配带宽;不能直接推广为音乐或音效质量指标 FAD 生成系统的分布级评价 样本量不足时严重虚高;绑定特征提取器;对单条样本没有意义 MOS(人工) 指定协议下的主观评价 有采样不确定性;跨批次比较需要协议、样本与评分锚点控制 UTMOS / DNSMOS MOS 的廉价代理 排序与绝对值都需校准,域外泛化尤其容易失败 两条总原则: 任何单一指标都可能给出与听感相反的结论。开头的教学场景说明,波形误差变小并不排除感知缺陷。多指标交叉 + 小规模人工抽听,是底线配置。 指标回答「像不像」,不回答「好不好」。一个模型可以无限逼近参考(各种距离都趋近 0),同时毫无创造力;生成类任务的最终评价一定要留人工的位置。 (具体缺陷怎么查、该看哪个数字,见下面「诊断手册:按症状选指标」。) 07. 诊断手册:按症状选指标与读图 前面几节讲的是「每个指标测什么」,这一节反过来——从你听到或收到的问题出发,倒推该看哪个数字。(它和上面的边界表是一体两面:表说「指标会骗你」,这里说「那该看谁」。) 7.1 先区分配对参考、参考分布与无参考预测 有对应参考波形时,可算 SI-SDR、SNR、谱距离以及适用的语音指标。没有逐条参考时,MOS 预测器仍可逐条打分,文本条件模型也可以计算文本与音频相似度。经典 FAD 是集合统计,不直接定位单条缺陷;研究中的单样本扩展有另外的定义与边界。 比较两个生成器时,优先固定 prompt 分布、后处理、时长和测试规模。若确实有一一对应的 prompt 与样本,可以做配对分析;非配对设计同样可以成立,但需要匹配设计的统计方法。报告效应量和置信区间,必要时使用置换检验或按独立音频单位重采样;多指标检验还要考虑多重比较。p=0.02 不能因为只有 24 条样本就改称「勉强」,p=0.2 也不能证明没有差异。单个集合级 FAD 更不能直接塞进逐条配对 t 检验。 7.2 按症状安排检查 持续高音、电流声:先看谱图确认是否存在额外窄带峰,再看该频带能量与逐帧平坦度。嘶声可能是宽带噪声,平坦度反而较高;蜂鸣、鸟叫、乐器泛音也可能是正常窄带成分。阈值必须由正常/异常验证集校准。 发闷或带宽不足:检查绝对电平、高频能量、频谱质心和相对参考的带宽。贝斯、鼓等低频内容天然窄带,不能由质心低直接判缺陷。 金属声或音乐噪声:沿时间看瞬态谱峰、带宽和谱距离变化,避免全段平均掩盖短时问题。有参考时可结合分离任务的伪影指标,再人工确认。 相位或通道异常:明确是否允许补偿全局延迟,检查通道间相关和对齐后的波形误差。低 SI-SDR 加小谱距离只能提示对齐问题,不能唯一证明根因。 整体增益偏差:用 RMS、LUFS 或参考增益比。SI-SDR 对单个全局非零增益不变,对随时间变化的增益并不不变;响度泵动需要逐窗观察。 可懂度不足:语音任务可结合 STOI 与人工转写;WER/CER 还依赖 ASR 本身,不能把识别器错误全部算成生成缺陷。 内容跑偏或多样性不足:在一致协议下看 FAD、条件一致性和多样性统计。特征方差变大也可能只是噪声,不能以「方差越大越好」替代覆盖度分析。 7.3 排查顺序和判定边界 先检查格式、采样率、声道、时长、绝对电平,再检查相对谱、局部异常和整体统计。接近静音时,按自身峰值归一化会把噪声放大得像正常内容,平坦度也可能很高;此时应单独报告静音比例及条件质量,不能默默删掉静音样本来提高总分。 浮点信号超过 0 dBFS 说明超出常见满刻度约定,但尚不能证明已发生硬削波。真峰值与采样峰值也要区分,是否压限应依据目标交付规范;不能假设流媒体平台的响度归一化会自动修复过载。THD+N 通常需要规定的测试信号,不是对任意音乐直接测 50/60 Hz 哼声的万能工具,后者应在谱上检查基频、谐波与背景的相对强度。 最后把可疑样本与 prompt、参考和听测对照:可视化说明信号发生了什么,不能单独决定这种变化是否符合内容要求。附录实验能验证增益、时移和新增纯音如何影响指标;实际产品的合格线需要自己的数据与听测来确定。 08. 经典论文脉络 PESQ(ITU-T P.862,2001):通信时代的遗产。把心理声学掩蔽模型塞进了客观指标,第一次让「机器预测听感」在窄带语音上有了可信度——也把「只适用语音」的边界刻进了骨子里。 SDR – half-baked or well done? (arXiv:1811.02508)(ICASSP 2019):指出 BSS_eval 那版 SDR 在单通道场景下会被滥用(改参考信号、靠调音量刷分),给出本文用的 SI-SDR 定义——误差与参考正交、不依赖幅度缩放。读它是理解「为什么偏偏叫 SDR」的最短路径。 Fréchet Audio Distance: A Metric for Evaluating Music Enhancement Algorithms (arXiv:1812.08466):把 FID 的思想搬到音频,提出基于背景集特征统计的音乐质量指标,解决了生成式音乐「没有正确答案可比」的问题;留下的问题是 FAD 数值绑定特征提取器、且强依赖样本量。 High Fidelity Neural Audio Compression (arXiv:2210.13438)(EnCodec):神经音频 codec 的代表作,结合客观指标与主观听测评价低码率重建。不能由这篇论文推导出「SI-SDR + PESQ + 梅尔距离」是所有 codec 的固定验收标准。 UTMOS: UTokyo-SaruLab System for VoiceMOS Challenge 2022 (arXiv:2204.02152):用 SSL 特征 + 集成学习预测 MOS,拿了 VoiceMOS Challenge 2022 多项第一。它代表的方向是「用模型替代昂贵的人工听测」——可行,但域外泛化仍是软肋。 一条暗线贯穿始终:指标的演进一直在追赶生成能力的演进。信号指标管不了感知,感知指标管不了分布,分布指标管不了单条样本——每一代生成技术把旧指标逼失效一次,然后逼出下一代指标。 09. 常见误解 「SI-SDR 高就是音质好」。它衡量去掉全局增益后的波形误差,对齐、相位、噪声和伪影都会影响它,但影响权重并非人耳感知权重。 「FAD 低说明每条音频都好」。均值与协方差相近不保证所有样本合格,也不保证更高阶分布一致、文本匹配或内容多样性。 「FAD 数字可以跨论文直接比」。特征器、参考集、时长、分帧、样本量、预处理一致,比较才有明确含义。本例是自定义梅尔特征距离,不应混称标准分数。 「MOS 不可复现」。MOS 是按协议收集的主观统计量,可以复做实验估计误差;跨批次须控制听测设计,报告不确定性。VoiceMOS 也评价 MSE 等指标,并非只比相关系数而完全不看绝对误差。 「无参考模型可替代人工验收」。MOS 预测器有域和数据偏差,即使在训练域也可能错排;发布标准需要经过人工标注或听测校准。 10. 动手验证 先安装 numpy matplotlib,把文末三个脚本存到同一目录,分别运行。它们不下载音频或权重;可复现的结论来自数值输出,不包括主观听感结论。 把 metric_disagreement.py 的 7000/7500 Hz 改成 200/250 Hz,重跑。高频占比应基本不变,但低频梅尔谱会改变。记录实际距离,不能预设「一定更小」或「一定同样刺耳」。这说明频带指标有明确的观测范围。 把 np.roll(x, 8) 改成 np.roll(x, 1),比较 SI-SDR 与谱距离。变化取决于波形自相关,不存在通用于所有音频的固定分数;真实延迟应在共同有效区间对齐,避免循环边界干扰。 在 FAD 脚本的循环中加入样本量 4,并多换几个种子。16 维协方差的秩最多为 3,估计很不稳定,但仍可计算;单次结果也可能偶然小。观察分布而非某一次抽样,才是样本量实验的正确读法。 11. 延伸阅读 FID / CLIP Score 到底测了什么(image_metrics):FAD 的直系前身。FID 的样本量陷阱、特征提取器绑定问题在图像侧一模一样,先读它再看本文会有「同一个故事换了个领域」的感觉。 视频生成评测:VBench 与人工验收(video_metrics):评测的维度从音频扩展到视频,多维度分解 + 人工验收的思路一脉相承。 离散化表征:VQ-VAE 与 VQGAN(vq_tokenizer):音频 codec(EnCodec、SoundStream)就是 VQ-VAE 思想在波形上的应用,想理解被评测的对象,从这篇进。 官方工具:fadtk(多种特征提取器的 FAD 工具箱)、torcheval 的 FrechetAudioDistance、ITU-T P.862.2(历史宽带 PESQ 与撤回说明)。 附录:完整代码 09 节用到的脚本全文如下(metric_disagreement.py、fad_sample_size.py、mel_spectrogram.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 metric_disagreement.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """同一个音频退化,四个指标给出四种排名。 这篇文章的核心论点就是这张表:指标不是「音质分数」,每个指标只回答一个 很具体的问题。用错了就会得到自欺欺人的结论。 三种退化分别是整体增益、窄带纯音和循环时移。 表中听感仅是待验证的假设,本文没有执行受试者听测。 干净的语音/音乐可以看成「基频 + 一串谐波」,这里就按这个造, 省得读者还要去找音频文件。只用 numpy。 """ import numpy as np SR = 16000 DUR = 1.0 N = int(SR * DUR) rng = np.random.default_rng(0) def make_clean(n=N, sr=SR): """基频 220Hz + 前 6 个谐波,加一点慢包络,听起来像有起伏的乐音。""" t = np.arange(n) / sr x = np.zeros(n) for k in range(1, 7): x += (1.0 / k) * np.sin(2 * np.pi * 220 * k * t) # 慢包络:2.5Hz 的轻微起伏,避免信号看起来像死板的稳态正弦 x *= 0.75 + 0.25 * np.sin(2 * np.pi * 2.5 * t) return x / np.abs(x).max() def degrade(x, kind, sr=SR): if kind == "A": return 0.95 * x # 音量差 5% if kind == "B": # 7/7.5 kHz 的窄带纯音:是否刺耳需听测,且两者均低于 Nyquist 频率 t = np.arange(len(x)) / sr buzz = 0.02 * ( np.sin(2 * np.pi * 7000 * t) + np.sin(2 * np.pi * 7500 * t) ) return x + buzz if kind == "C": return np.roll(x, 8) # 时移 8 采样 = 0.5ms raise ValueError(kind) def si_sdr(ref, est): """Scale-Invariant SDR。 先把 est 投影到 ref 上(找一个最优缩放 alpha),剩下的残差才算失真。 所以「整体音量变了」在 SI-SDR 里被完全抵消——这是设计意图, 但也意味着它测不出增益错误。 """ if np.dot(ref, ref) <= 0 or np.dot(est, est) <= 0: raise ValueError("SI-SDR needs non-silent reference and estimate") alpha = np.dot(est, ref) / np.dot(ref, ref) e_target = alpha * ref e_res = est - e_target num = np.dot(e_target, e_target) den = np.dot(e_res, e_res) if den <= 1e-24 * num: return float("inf"), alpha return 10 * np.log10(num / den), alpha def snr(ref, est): """不做缩放对齐的信噪比。音量差会被算成失真。""" res = est - ref den = np.dot(res, res) if den == 0: return float("inf") return 10 * np.log10(np.dot(ref, ref) / den) def stft_mag(x, n_fft=512, hop=128): win = np.hanning(n_fft) n_frames = 1 + (len(x) - n_fft) // hop frames = np.stack([x[i * hop:i * hop + n_fft] * win for i in range(n_frames)]) return np.abs(np.fft.rfft(frames, n_fft, axis=1)).T # [freq, time] def hz_to_mel(f): return 2595.0 * np.log10(1.0 + f / 700.0) def mel_to_hz(m): return 700.0 * (10 ** (m / 2595.0) - 1.0) def mel_filterbank(sr=SR, n_fft=512, n_mels=64, fmin=20.0, fmax=8000.0): n_freqs = n_fft // 2 + 1 hz = mel_to_hz(np.linspace(hz_to_mel(fmin), hz_to_mel(fmax), n_mels + 2)) bins = np.floor(hz * n_fft / sr).astype(int) fb = np.zeros((n_mels, n_freqs)) for i in range(n_mels): left, center, right = bins[i], bins[i + 1], bins[i + 2] if center <= left: center = left + 1 if right <= center: right = center + 1 for k in range(left, min(center, n_freqs)): fb[i, k] = (k - left) / (center - left) for k in range(center, min(right, n_freqs)): fb[i, k] = (right - k) / (right - center) return fb def log_mel(x, sr=SR, n_fft=512, hop=128, n_mels=64): """归一化梅尔幅度谱的自然对数,单位为 ln 幅度(不是 dB)。 每条样本独立归一化是本实验为了忽略整体增益所作的选择, 不是所有音频评测的必需步骤。幅度 floor=1e-4 对应峰值下 80 dB。 """ fb = mel_filterbank(sr=sr, n_fft=n_fft, n_mels=n_mels) mel = fb @ stft_mag(x, n_fft=n_fft, hop=hop) mel = mel / (mel.max() + 1e-12) return np.log(np.maximum(mel, 1e-4)) # floor 在 log(1e-4) ≈ -9.2 def log_mel_distance(ref, est): """对数梅尔谱的平均绝对误差(单位为自然对数幅度差)。""" return float(np.mean(np.abs(log_mel(est) - log_mel(ref)))) def hf_energy_ratio(x, sr=SR, n_fft=512, hop=128, f_lo=6000.0): """6kHz 以上能量占总能量的比例——高频毛刺会把它顶起来。""" mag = stft_mag(x, n_fft=n_fft, hop=hop) freqs = np.fft.rfftfreq(n_fft, 1.0 / sr) power = (mag ** 2).mean(axis=1) return float(power[freqs >= f_lo].sum() / power.sum()) def spectral_flatness(x, sr=SR, n_fft=512, hop=128, f_lo=4000.0, f_hi=9000.0): """高频平均功率谱的几何均值/算术均值,范围 0~1。 小值表示谱能量集中,不等价于存在缺陷;纯乐音也会很低。 只在该频带有足够能量时解释,阈值必须按任务和频带校准。 与高频能量占比一样,理想计算对非零全局增益不变。 """ mag = stft_mag(x, n_fft=n_fft, hop=hop) freqs = np.fft.rfftfreq(n_fft, 1.0 / sr) band = mag[(freqs >= f_lo) & (freqs <= f_hi)] power = (band ** 2).mean(axis=1) + 1e-20 geo = np.exp(np.mean(np.log(power))) ari = np.mean(power) return float(geo / ari) def main(): ref = make_clean() print(f"参考信号: {len(ref)} 采样 = {len(ref) / SR:.2f}s @ {SR}Hz") print(f"参考信号高频能量占比: {hf_energy_ratio(ref):.3e}") print() header = (f"{'退化':<18}{'操作':<12}{'SI-SDR(dB)':>12}{'SNR(dB)':>10}" f"{'梅尔距离':>10}{'高频占比':>12}{'高频谱平坦度':>14}") print(header) print("-" * len(header)) rows = {} for kind, feel in (("A", "增益变化"), ("B", "新增纯音"), ("C", "循环时移")): est = degrade(ref, kind) s, alpha = si_sdr(ref, est) d = { "si_sdr": s, "snr": snr(ref, est), "mel": log_mel_distance(ref, est), "hf": hf_energy_ratio(est), "flat": spectral_flatness(est), "alpha": alpha, } rows[kind] = d print(f"{kind + ' ' + {'A': '音量×0.95', 'B': '高频毛刺', 'C': '时移0.5ms'}[kind]:<18}" f"{feel:<12}{s:>12.2f}{d['snr']:>10.2f}{d['mel']:>10.4f}" f"{d['hf']:>12.3e}{d['flat']:>14.4f}") print() print("读数(这才是重点):") print(f" A 音量小 5%:SI-SDR = {rows['A']['si_sdr']:.2f} dB(缩放被抵消,几乎满分)," f"SNR = {rows['A']['snr']:.2f} dB(老老实实记了一笔)") print(f" → 最优缩放 alpha = {rows['A']['alpha']:.4f},表示估计波形为参考的 0.95 倍") print(f" B 高频毛刺:SI-SDR = {rows['B']['si_sdr']:.2f} dB 看着还行," f"高频占比却从 {hf_energy_ratio(ref):.3e} 涨到 {rows['B']['hf']:.3e}") print(f" C 时移 0.5ms:SI-SDR 掉到 {rows['C']['si_sdr']:.2f} dB," f"但梅尔距离只有 {rows['C']['mel']:.4f} —— 谱几乎没动") print() print("结论:SI-SDR 衡量的是「波形对不对齐」,梅尔距离衡量的是「频谱像不像」,") print(" 高频占比描述频带能量分配,是否刺耳还需要听测。") if __name__ == "__main__": main() fad_sample_size.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """FAD 到底在比什么,以及为什么样本量不够时它根本不可信。 FAD(Fréchet Audio Distance)是 FID 的音频版:不问「这两条音频像不像」, 而是问「这两堆音频的分布像不像」。这个区别决定了它有两个反直觉的性质: 1. 它不配对。不需要参考音频和测试音频一一对应, 所以能评测「凭空生成」的音乐/音效——这是 PESQ 之类做不到的。 2. 它对样本量极其敏感。协方差是从样本里估的,样本不够时估出来的协方差 本身就是歪的,算出来的 FAD 跟真实值差得离谱。 这里用 16 维归一化对数梅尔幅度均值当特征(真正的 FAD 用 VGGish 的 128 维 embedding, 我们没有那个模型,但 Fréchet 距离这一步的数学完全一样)。 """ import numpy as np SR = 16000 N_FFT = 512 HOP = 128 N_MELS = 16 rng = np.random.default_rng(7) def stft_mag(x, n_fft=N_FFT, hop=HOP): win = np.hanning(n_fft) n_frames = 1 + (len(x) - n_fft) // hop frames = np.stack([x[i * hop:i * hop + n_fft] * win for i in range(n_frames)]) return np.abs(np.fft.rfft(frames, n_fft, axis=1)).T def hz_to_mel(f): return 2595.0 * np.log10(1.0 + f / 700.0) def mel_to_hz(m): return 700.0 * (10 ** (m / 2595.0) - 1.0) def mel_filterbank(sr=SR, n_fft=N_FFT, n_mels=N_MELS, fmin=20.0, fmax=8000.0): n_freqs = n_fft // 2 + 1 hz = mel_to_hz(np.linspace(hz_to_mel(fmin), hz_to_mel(fmax), n_mels + 2)) bins = np.floor(hz * n_fft / sr).astype(int) fb = np.zeros((n_mels, n_freqs)) for i in range(n_mels): left, center, right = bins[i], bins[i + 1], bins[i + 2] if center <= left: center = left + 1 if right <= center: right = center + 1 for k in range(left, min(center, n_freqs)): fb[i, k] = (k - left) / (center - left) for k in range(center, min(right, n_freqs)): fb[i, k] = (right - k) / (right - center) return fb FB = mel_filterbank() def embed(x): """一段音频 → 16 维特征:各梅尔带的平均对数幅度。""" mel = FB @ stft_mag(x) mel = mel / (mel.max() + 1e-12) return np.log(np.maximum(mel, 1e-4)).mean(axis=1) def make_tone(f0, n=8000, buzz=False, snr_db=None): """一个乐音:基频 f0 + 谐波。buzz=True 时叠高频毛刺。""" t = np.arange(n) / SR x = np.zeros(n) for k in range(1, 6): x += (1.0 / k) * np.sin(2 * np.pi * f0 * k * t) if buzz: x += 0.05 * (np.sin(2 * np.pi * 7000 * t) + np.sin(2 * np.pi * 7500 * t)) if snr_db is not None: x = x + rng.normal(0, np.sqrt(np.mean(x ** 2) / (10 ** (snr_db / 10))), n) return x / (np.abs(x).max() + 1e-12) def sqrtm_psd(A): """对称半正定阵的平方根;仅裁掉舍入级负特征值。""" A = (A + A.T) / 2 w, V = np.linalg.eigh(A) if w.min() < -1e-10 * max(1.0, np.abs(w).max()): raise ValueError("matrix is not positive semidefinite") w = np.clip(w, 0.0, None) return (V * np.sqrt(w)) @ V.T def frechet_distance(X, Y): """两个特征集合的 Fréchet 距离。 把每堆特征当成一个高斯(只关心均值和协方差),再算这两个高斯之间的 2-Wasserstein 距离的平方。闭式解: FD = ||mu_x - mu_y||^2 + Tr(S_x + S_y - 2 (S_x S_y)^{1/2}) """ mu_x, mu_y = X.mean(axis=0), Y.mean(axis=0) S_x = np.cov(X, rowvar=False) S_y = np.cov(Y, rowvar=False) diff = mu_x - mu_y # 不可把一般非对称的 S_x @ S_y 直接交给 eigh。 # 使用对称半正定夹心矩阵,平方根的迹与原式相同。 s_x_sqrt = sqrtm_psd(S_x) covmean = sqrtm_psd(s_x_sqrt @ S_y @ s_x_sqrt) fd = diff @ diff + np.trace(S_x) + np.trace(S_y) - 2 * np.trace(covmean) tol = 1e-7 * max(1.0, np.trace(S_x) + np.trace(S_y)) if fd < -tol: raise ArithmeticError("negative distance exceeds numerical tolerance") return float(max(fd, 0.0)) # 仅容差内的负舍入误差截到 0 def main(): # 参考集:干净乐音,基频在 200~400Hz 之间随机 # 测试集 A:同分布(只是另一批采样)→ FAD 应该接近 0 # 测试集 B:带高频毛刺 → FAD 应该明显大于 0 def make_set(n, buzz=False): return np.stack([ embed(make_tone(rng.uniform(200, 400), buzz=buzz)) for _ in range(n) ]) print("特征维度: 16(各梅尔带的平均对数幅度)") print() print("样本量对 FAD 的影响(同一个分布对,只是抽的样本数不同):") print(f"{'每堆样本数':>10}{'同分布 FAD':>14}{'带毛刺 FAD':>14}{'相对基线':>10}") print("-" * 50) big_ref = make_set(600) for n in (8, 16, 32, 64, 128, 300): ref = big_ref[:n] same = np.stack([ embed(make_tone(rng.uniform(200, 400))) for _ in range(n) ]) buzz = make_set(n, buzz=True) fd_same = frechet_distance(ref, same) fd_buzz = frechet_distance(ref, buzz) ok = f"{fd_buzz / max(fd_same, 1e-6):.1f}x" print(f"{n:>10}{fd_same:>14.4f}{fd_buzz:>14.4f}{ok:>10}") print() print("注意最上面几行:样本只有 8 或 16 个时,「同分布」的 FAD 就已经不是 0 了,") print("这是均值与协方差估计误差;16 维、16 个样本的中心化协方差必然奇异。") print("真实 VGGish 特征为 128 维,但所需样本数还取决于谱结构、独立性和精度——") print("这就是为什么论文里报的 FAD 必须同时报样本数,否则数字不可比。") print() print("另一面:") print(f" 300 个样本时,同分布 FAD 与带毛刺 FAD 差了 " f"{frechet_distance(big_ref[:300], make_set(300, buzz=True)) / max(frechet_distance(big_ref[:300], np.stack([embed(make_tone(rng.uniform(200, 400))) for _ in range(300)])), 1e-9):.1f} 倍") print(" → 倍数只是描述量,不是统计检验;需多次抽样或按音频做 bootstrap。") if __name__ == "__main__": main() mel_spectrogram.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """把「指标打架」画出来看:三种退化的梅尔谱。 波形指标(SI-SDR)分不清的三件事,在谱上一眼就能看出来: 音量 ×0.95 —— 本图每条单独归一化,整体增益被抵消 高频毛刺 —— 谱的中上部多出一条亮带(听感需另行验证) 时移 0.5ms —— 几乎看不出差别(但 SI-SDR 可能明显下降) 产出 figures/mel_compare.png,文章里用  引用。 """ import os import numpy as np import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt from matplotlib import font_manager # matplotlib 默认字体没有中文,标题会渲染成一排豆腐块。 # 按 macOS / Linux 常见字体逐个试,找一个系统里真有的。 _AVAILABLE = {f.name for f in font_manager.fontManager.ttflist} for _cand in ["PingFang SC", "Heiti SC", "Arial Unicode MS", "Songti SC", "STHeiti", "Noto Sans CJK SC", "WenQuanYi Zen Hei"]: if _cand in _AVAILABLE: plt.rcParams["font.sans-serif"] = [_cand] break plt.rcParams["axes.unicode_minus"] = False # 负号不要用中文缺字替代 SR = 16000 N_FFT = 512 HOP = 128 N_MELS = 64 OUT_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "figures") def make_clean(n=16000, sr=SR): t = np.arange(n) / sr x = np.zeros(n) for k in range(1, 7): x += (1.0 / k) * np.sin(2 * np.pi * 220 * k * t) x *= 0.75 + 0.25 * np.sin(2 * np.pi * 2.5 * t) return x / np.abs(x).max() def degrade(x, kind, sr=SR): t = np.arange(len(x)) / sr if kind == "A": return 0.95 * x if kind == "B": return x + 0.02 * (np.sin(2 * np.pi * 7000 * t) + np.sin(2 * np.pi * 7500 * t)) if kind == "C": return np.roll(x, 8) raise ValueError(kind) def stft_mag(x, n_fft=N_FFT, hop=HOP): win = np.hanning(n_fft) n_frames = 1 + (len(x) - n_fft) // hop frames = np.stack([x[i * hop:i * hop + n_fft] * win for i in range(n_frames)]) return np.abs(np.fft.rfft(frames, n_fft, axis=1)).T def hz_to_mel(f): return 2595.0 * np.log10(1.0 + f / 700.0) def mel_to_hz(m): return 700.0 * (10 ** (m / 2595.0) - 1.0) def mel_filterbank(sr=SR, n_fft=N_FFT, n_mels=N_MELS, fmin=20.0, fmax=8000.0): n_freqs = n_fft // 2 + 1 hz = mel_to_hz(np.linspace(hz_to_mel(fmin), hz_to_mel(fmax), n_mels + 2)) bins = np.floor(hz * n_fft / sr).astype(int) fb = np.zeros((n_mels, n_freqs)) for i in range(n_mels): left, center, right = bins[i], bins[i + 1], bins[i + 2] if center <= left: center = left + 1 if right <= center: right = center + 1 for k in range(left, min(center, n_freqs)): fb[i, k] = (k - left) / (center - left) for k in range(center, min(right, n_freqs)): fb[i, k] = (right - k) / (right - center) return fb FB = mel_filterbank() def log_mel_db(x): mel = FB @ stft_mag(x) mel = mel / (mel.max() + 1e-12) return 20 * np.log10(np.maximum(mel, 1e-4)) # 归一化后取 dB,floor 在 -80dB def main(): ref = make_clean() cases = [ ("参考(干净)", ref), ("A 音量 ×0.95", degrade(ref, "A")), ("B 高频毛刺", degrade(ref, "B")), ("C 时移 0.5ms", degrade(ref, "C")), ] specs = [log_mel_db(x) for _, x in cases] # 每条谱已独立按峰值归一化,共用色标只用于比较相对谱形状。 vmin = min(s.min() for s in specs) vmax = max(s.max() for s in specs) fig, axes = plt.subplots(1, 4, figsize=(15, 3.6), constrained_layout=True) for ax, (name, _), s in zip(axes, cases, specs): im = ax.imshow(s, origin="lower", aspect="auto", vmin=vmin, vmax=vmax, cmap="magma") ax.set_title(name, fontsize=12) ax.set_xlabel("帧") ax.set_ylabel("梅尔带") fig.colorbar(im, ax=axes, shrink=0.8, label="dB(相对峰值)") fig.suptitle("同一个音频的四种形态:梅尔谱", fontsize=13) os.makedirs(OUT_DIR, exist_ok=True) out = os.path.join(OUT_DIR, "mel_compare.png") fig.savefig(out, dpi=150, bbox_inches="tight", facecolor="white") plt.close(fig) print(f"图已保存: {out}") print() print("每张谱的动态范围与「和参考谱的平均绝对差」:") print(f"{'形态':<16}{'最小值(dB)':>12}{'最大值(dB)':>12}{'与参考的平均差(dB)':>20}") print("-" * 62) base = specs[0] for (name, _), s in zip(cases, specs): print(f"{name:<16}{s.min():>12.2f}{s.max():>12.2f}" f"{np.mean(np.abs(s - base)):>20.3f}") if __name__ == "__main__": main()
2026年09月21日
7 阅读
0 评论
2 点赞
2026-09-20
AIGC 基本功|视频 VAE 的常见 loss 组合-VAELoss
视频 VAE 的常见 loss 组合 所属方向:表征与压缩 | 难度:工程实战 | 前置知识:视频 VAE 的时空压缩结构(本文第 02 节另补最小背景) 关键词:重建损失、KL 损失、LPIPS、感知损失、GAN 损失、判别器、loss 权重 01. 为什么需要它 训练一个视频 VAE,最朴素的想法是「编码器压缩、解码器还原,加个 L1 或 KL 就够了」。不同目标可能暴露不同问题,效果取决于瓶颈、数据和训练配方: 只用 L1/L2 + KL:画面是稳的,但像蒙了一层猪油 —— 头发丝、文字、树叶这种高频纹理全被抹掉,人脸带着一层「塑料感」。在存在重建不确定性时,这与逐像素回归的性质有关,但不能断言 L1/L2 + KL 必然不可用(L2 的最优解是条件期望,天然倾向平均、发糊)。 加上 GAN 想救清晰度:判别器一旦过早发力或过强,解码器立刻走样 —— 出现棋盘格、闪烁、甚至凭空捏造纹理;训练损失曲线看起来在降,重建却越来越假。 把图像那套搬到视频上:每一帧单独看 PSNR 都不错,连起来播放却疯狂闪烁,相邻帧的纹理在抖动。单帧指标根本测不出时间维的失真。 这不是三个孤立的 bug,而是三类不同的失真,需要三种不同的尺子去量。视频 VAE 的训练目标因此通常是四项的组合:像素重建、KL 正则、感知损失、对抗损失,有时再加一项时间一致性约束。玄学的地方从来不是「用哪几项」—— 这几项高度趋同 —— 而是权重配比:为什么 HunyuanVideo 的 KL 权重能小到 $10^{-6}$,而感知项是 $0.1$、对抗项是 $0.05$?这些数字不是拍脑袋,是被各项的量纲、归一化口径和训练阶段逼出来的。 本文把每一项「在管什么失真、数学上长什么样、量级有多大、什么时候帮倒忙」逐一拆开,所有关键数字都来自随文可跑的脚本(只依赖 numpy)。 02. 最小可用理解 最小背景(前置 VideoVAE):视频 VAE 用 3D 因果卷积把一段视频 $x\in\mathbb{R}^{B\times C\times T\times H\times W}$ 压成低维潜变量 $z$(典型压缩率:时间 $4\times$、空间 $8\times$),解码器再从 $z$ 重建出 $r$。编码器输出一个高斯后验 $q(z\mid x)=\mathcal{N}(\mu,\sigma^2)$,训练目标要同时满足三件事:重建要像、潜变量分布要规整(好让后面的扩散模型去建模)、压缩率要够高。loss 就是在这三者之间做权衡的旋钮。 三句话讲清四项损失: 像素损失(L1/L2)和 KL 是「保正确」的:前者逐像素对齐内容与结构,后者把潜变量摁在标准正态附近、防止它为了重建而无限膨胀。 LPIPS 感知损失和 GAN 对抗损失是「保好看」的:LPIPS 在预训练深度网络的特征空间里比较,惩罚人眼在意的语义 / 纹理差异;GAN 让一个判别器去挑刺,逼着解码器还原高频细节。 这四项量纲完全不同,必须加权配平,GAN 常用「晚启动 + 自适应权重」来控制训练平衡:先让重建项把解码器教到成形,再让判别器入场,并按两项在最后一层上的梯度范数之比动态调权。视频还要额外盯时间维(时空判别器或帧间一致性项)。 03. 数学推导 3.1 从 ELBO 到「重建 + KL」两项 VAE 最大化证据下界(ELBO)。写成「最小化负 ELBO」,目标天然裂成两项: $$\mathcal{L} = \underbrace{-\mathbb{E}_{q_\phi(z\mid x)}[\log p_\theta(x\mid z)]}_{\text{重建项}} + \underbrace{D_{\mathrm{KL}}\big(q_\phi(z\mid x)\,\|\,p(z)\big)}_{\text{KL 正则项}}$$ 逐项解释符号:$\phi$ 是编码器参数、$\theta$ 是解码器参数;$q_\phi(z\mid x)$ 是编码器给出的后验分布;$p_\theta(x\mid z)$ 是解码器的似然;$p(z)=\mathcal{N}(0,I)$ 是标准正态先验。重建项要求「从采样出的 $z$ 能还原 $x$」,KL 项要求「后验别离先验太远」—— 它是潜变量的正则项,鼓励与先验对齐;它不保证感知空间的平滑或连通,也不是所有 latent 扩散能够学习的必要条件。 当后验取对角高斯 $q_\phi(z\mid x)=\mathcal{N}(\mu_\phi(x),\,\mathrm{diag}(\sigma_\phi(x)^2))$、先验取标准正态时,KL 有解析解(不用采样、不用估计)。令 $\ell_j=\log\sigma_j^2$(工程上网络直接预测 $\ell$,数值更稳): $$D_{\mathrm{KL}} = -\frac{1}{2}\sum_{j=1}^{d_z}\Big(1+\ell_j-\mu_j^2-\exp(\ell_j)\Big)$$ 这里求和下标 $j$ 跑遍所有潜变量维度。直觉:$\tfrac12\mu_j^2$ 惩罚均值偏离 0,$-(1+\ell_j-e^{\ell_j})$ 惩罚方差偏离 1($\ell=0$ 即 $\sigma^2=1$ 时该项为 0)。 第一个容易被忽略的坑是「求和 vs 平均」。教科书公式对 $d_z$ 个维度求和;而像素损失通常对全部像素求平均。两者元素个数差着几个数量级,直接相加 KL 会凭「项数多」碾压重建。随文脚本 vaeloss_terms.py 在一个合成视频上实测: # 解析 KL,logvar 是网络直接预测的 log(sigma^2) kl_per_element = -0.5 * (1 + logvar - mu ** 2 - np.exp(logvar)) kl_sum, kl_mean = kl_per_element.sum(), kl_per_element.mean() l1 = np.mean(np.abs(x - r)) print("L1 = %.5f" % l1) # 0.03158 print("KL(sum) = %.3f KL(mean) = %.4f" % (kl_sum, kl_mean)) print("KL(sum)/L1 = %.0f 倍" % (kl_sum / l1)) 输出(视频 $2\times3\times8\times32\times32$,latent $2\times4\times2\times4\times4$,像素 49152 个、全 batch 的 latent 共 256 个数值(每样本 128 个)): L1 重建 (mean) 0.03158 KL (全 batch sum,教学口径) 57.66490 KL (mean,换口径) 0.22525 KL(sum)/L1 的量级倍数 : 1826.0x 归约口径是开源视频 VAE 的 KL 权重有时出现 $10^{-6}$ 这种「看着离谱」的数字 —— 的原因之一,但不能由这个 toy 直接反推实际系数。通常应先对每样本 latent 求和,再对 batch 平均;这里为演示求和效应,57.66 是整个 batch 的和,batch=2 时常见口径为 28.83。抄权重之前先对齐归一化口径,这句话后面还会强调。 这张图要看什么:左图是四项损失在同一个合成视频上的原始数值(对数轴)——KL(sum) 一根柱子顶到 57.66,是 L1 的 1826 倍;右图套上 HunyuanVideo 的配方权重之后,各项贡献变为不同数量级,不能称为完全配平,最矮的 KL 贡献只有 $5.8\times10^{-5}$。权重同时反映口径、梯度和重建目标,不能仅按标量 loss 大小配平(图由 code/make_figures.py 生成,下同)。 3.2 像素损失:为什么多用 L1 而不是 L2 $$\mathcal{L}_{\mathrm{pix}} = \frac{1}{N}\sum_{n=1}^{N}\big\|x_n-r_n\big\|_1$$ L2(MSE)对大误差平方加权,其逐点最优解是给定 $z$ 下所有可能输出的条件均值—— 面对「这块纹理可能是 A 也可能是 B」的不确定性,它选择把 A、B 平均掉,结果就是模糊。L1 在非零残差处的梯度为 $\pm1$,其逐点最优解是条件中位数,不会因为误差大就给出更强的「往平均靠」的驱动力,对异常值通常更稳健,但条件中位数也可能模糊,不能保证所有纹理更锐利。代价是 L1 在零点不可导、对小误差的梯度恒定,容易留下颗粒感 —— 这正是要靠感知 / 对抗项补的地方。也有工作(如 LTX-Video)在像素项里混用 MSE 与小波域 L1(Video-DWT),在多尺度上约束。 3.3 感知损失 LPIPS:换一把「人眼的尺子」 像素距离有个致命问题:同样大小的像素误差,人眼感受天差地别。一个全局亮度偏移,MSE 不小但人眼几乎无感;一团等量的逐点噪声,MSE 相同却把纹理毁了。LPIPS(Learned Perceptual Image Patch Similarity)改用在 ImageNet 上预训练的分类网络(VGG/Alex,参数冻结)提特征,再在特征空间量距离: $$\mathcal{L}_{\mathrm{LPIPS}}(x,r) = \sum_{k}\frac{1}{H_kW_k}\Big\|\,w_k\odot\big(\hat{y}_x^{k}-\hat{y}_r^{k}\big)\Big\|_2^2$$ 符号:$y_x^k$、$y_r^k$ 是第 $k$ 层(LPIPS 用 VGG 的 relu1_2 到 relu5_3 共 5 层)对真实图和重建图提的特征图;$\hat{y}$ 表示沿通道做了归一化(除以通道维 L2 范数);公式中的 $w_k^2$ 对应实现里作用在平方特征差上的 $1\times1$ 非负通道权重;$w_k$ 可理解为它的平方根,而非直接把卷积权重再平方。该权重在人类主观偏好数据集 BAPPS 上学出来、之后冻结;最后空间平均、跨层求和。两个细节缺一不可:逐通道归一化让比较不被某些高幅值通道主导,学习权重 $w_k$ 让「哪些层的差异人眼更在意」由数据决定。它天然偏向语义 / 结构 / 纹理,而对整体明暗、轻微色偏不敏感。 pixel_vs_perceptual.py 用一个免权重的多尺度高通特征(高斯差分 DoG)复现 LPIPS 的结构,对比三种退化: # A:全局亮度偏移;B:同等 L2 能量的逐点噪声;C:高斯模糊 rA = x + delta # delta = 0.12 rB = x + rng.normal(0, delta, x.shape) # 标准差同为 0.12 # 像素 MSE:A 约等于 B;特征距离:B 远大于 A 真实输出: 重建方式 像素MSE 感知距离 感知/像素 A 亮度偏移(整体+0.12) 0.01440 0.00008 0.01 B 逐点噪声(σ=0.12) 0.01426 0.64084 44.93 C 高斯模糊 0.00508 0.50163 98.69 A、B 的像素 MSE 只差 0.95%,但特征距离里 B 是 A 的约 $8\times10^3$ 倍;模糊 C 的像素误差最小,感知距离却很高。高通教学代理天然抑制直流亮度,因此本例差距尤其大;真实 LPIPS 使用学习特征和权重,不能据此断言它对亮度或颜色变化不敏感。 3.4 对抗损失:雇一个判别器专挑高频毛病 GAN 引入一个判别器 $D$,训练它区分真实帧与重建帧;解码器(生成器)则努力骗过它。视频 / 图像重建里几乎都用 PatchGAN 式判别器:不输出整图真假,而是对每个空间 patch 打分,主要约束其感受野内的统计,常改善局部纹理,但并非数学上只看高频。最常用的 hinge 形式: $$\mathcal{L}_{D} = \tfrac{1}{2}\Big(\mathbb{E}_{\text{real}}[\max(0,\,1-D(x))]+\mathbb{E}_{\text{recon}}[\max(0,\,1+D(r))]\Big)$$ $$\mathcal{L}_{G}^{\mathrm{adv}} = -\mathbb{E}_{r}[D(r)]$$ 判别器希望真帧打分大于 1、重建打分小于 -1;生成器希望重建打分尽量大(为正)。gan_schedule.py 用合成 logits 算了判别器在三种强弱下的损失: 起步:判别器分不清 D(hinge)=1.0126 G(hinge)=-0.0312 real≈ 0.01 fake≈ 0.03 健康:适度拉开 D(hinge)=0.0767 G(hinge)= 1.2067 real≈ 1.21 fake≈-1.21 过强:margin 已饱和 D(hinge)=0.0000 G(hinge)= 5.9713 real≈ 6.02 fake≈-5.97 注意判别器 hinge loss 为 0 只说明这些样本的 margin 已满足,此时判别器对应损失的梯度为 0;不意味着生成器梯度消失。这里 $\mathcal L_G=-\mathbb E[D(r)]$ 对 fake logit 的导数仍是 −1,传回解码器的梯度还取决于判别器对输入的导数。仅看 +6/−6 logits 无法判断梯度是否健康。 其一,GAN 晚启动(warm-up)。重建项还没把解码器教出基本形状时,判别器挑的「毛病」没有意义,甚至会把训练带偏。taming 的做法是一个开关 adopt_weight,在第 $t_{\mathrm{start}}$ 步前把对抗权重置 0: $$\delta(t)=\begin{cases}0,&t<t_{\mathrm{start}}\\ \lambda_{\mathrm{adv}},&t\ge t_{\mathrm{start}}\end{cases}$$ 脚本实测 disc_start=2000 时,step 0 和 1999 的权重为 0,step 2000 起跳到 0.05。 其二,自适应对抗权重。固定 $\lambda_{\mathrm{adv}}$ 的两难:训练早期重建梯度很大、GAN 抢不过;后期重建收敛、梯度变小,同样的 GAN 梯度又会相对越来越强甚至压过重建。VQGAN 的解法是让两项在解码器最后一层权重 $w_L$ 上的梯度范数之比来决定权重: $$\lambda_{\mathrm{adv}}(t)=\mathrm{clip}\left(\frac{\|\nabla_{w_L}\mathcal{L}_{\mathrm{rec}}\|}{\|\nabla_{w_L}\mathcal{L}_{G}^{\mathrm{adv}}\|+\epsilon},\ 0,\ 10^4\right)\cdot\delta(t)$$ 直觉:它动态地让「对抗项在最后一层上产生的梯度量级」与「重建项的梯度量级」匹配,重建没收敛时比值大、收敛后比值自动变小。adaptive_weight.py 用一个可解析求导的迷你线性解码器精确计算两个梯度范数,并用中心差分验证(解析 0.012371 对差分 0.012371;0.255913 对 0.255913): 训练早期(未收敛) L_rec=0.73740 ||∇rec||=0.4071 ||∇gan||=4.2840 d_weight=0.0950 训练后期(近收敛) L_rec=0.00071 ||∇rec||=0.0127 ||∇gan||=4.2840 d_weight=0.0030 重建收敛后自适应权重从 0.095 掉到 0.003(约 32 倍),GAN 项被自动调小 —— 这正是固定权重给不了的能力。 这张图要看什么:横轴是合成出来的训练进程(从「解码器还没学会」到「重建已收敛」),三条曲线分别是重建损失、重建项在最后一层上的梯度范数、以及据此算出的 $\lambda_{\mathrm{adv}}$。要点是 $\lambda_{\mathrm{adv}}$ 不是人为排的衰减计划,而是被 $\|\nabla\mathcal{L}_{\mathrm{rec}}\|$ 拖着走的:早期 0.095、后期 0.003。对照上面的 warm-up 开关看更清楚——$\delta(t)$ 负责「第 2000 步之前完全不启用」,这条曲线负责「启用之后给多大」,两者相乘才是最终权重。 3.5 总目标,以及视频多出的时间维 把四项合起来,连续潜变量视频 VAE 的生成器目标是: $$\mathcal{L}_{G} = \lambda_{\mathrm{pix}}\mathcal{L}_{\mathrm{pix}}+\lambda_{\mathrm{p}}\mathcal{L}_{\mathrm{LPIPS}}+\lambda_{\mathrm{kl}}\mathcal{L}_{\mathrm{KL}}+\lambda_{\mathrm{adv}}(t)\,\mathcal{L}_{G}^{\mathrm{adv}}$$ (VQGAN 是离散码本,没有 KL,对应位置换成 codebook/commitment 损失;连续 KL-VAE 才是上面这版。)判别器另用 $\mathcal{L}_D$ 单独更新,二者交替。 视频比图像多一维,单帧损失管不到帧间。temporal_consistency.py 构造了一段运动视频,给两种重建:A 加逐帧独立噪声(播放时闪烁),B 加跨帧恒定的退化(不闪),两者单帧空间 L1 几乎相同: 重建 单帧空间L1 时序差分L1 时序差分MSE A 逐帧独立噪声(闪) 0.06393 0.09018 0.01275 B 跨帧恒定退化(稳) 0.06340 0.00455 0.00003 单帧 L1 几乎相等(0.0639 对 0.0634),时序差分 MSE 却差了约 393 倍。 这张图要看什么:上半部是同一个像素点随帧走的亮度轨迹。A(逐帧独立噪声)和 B(跨帧恒定退化)的逐帧平均误差几乎一样,但 A 的轨迹在真值附近高频锯齿抖动,B 的轨迹只是整体平移——前者播放起来就是闪烁,后者只是画质差一点。下半部把这件事量化:单帧空间 L1 两根柱子几乎齐平(0.0639 / 0.0634),时序差分 L1 差 20 倍,时序差分 MSE 差到 393 倍。所以「视频 VAE 的重建指标好看但看着闪」,根因是指标选错了:空间指标对时间频率不敏感。 补救有两条路:一是把判别器从 2D 扩到时空(3D/PatchGAN、混合外观 - 运动判别器),让它直接看片段;二是显式加时间一致性损失,用光流把相邻帧 warp 过来再比(静止区域退化成帧差): $$\mathcal{L}_{\mathrm{temp}}=\frac{1}{T-1}\sum_{t=1}^{T-1}\big\|r_t-\mathcal{W}(r_{t-1},\,F_{t\to t-1})\big\|_1$$ 其中 $\mathcal{W}$ 是 warp 算子、$F$ 是估计的光流;没有光流时可用相邻帧差分近似 $\mathcal{L}_{\Delta}=\frac{1}{T-1}\sum_t\|(r_t-r_{t-1})-(x_t-x_{t-1})\|_1$。 04. 代码实现 随文 5 个脚本只依赖 numpy,python 脚本名.py 即可运行,完整版在文末附录;另有 make_figures.py 需要 matplotlib,负责本文三张配图。这里串起主干。 (1)四项损失与量级对照(vaeloss_terms.py):在合成视频上算 L1、解析 KL(sum/mean 两种口径)、感知代理、GAN,核心是 KL 解析式: kl_per_element = -0.5 * (1 + logvar - mu ** 2 - np.exp(logvar)) kl_sum = kl_per_element.sum() # 教科书口径:对 latent 维求和 l1 = np.mean(np.abs(x - r)) # 工程口径:对像素求平均 print("KL(sum)/L1 = %.0f 倍" % (kl_sum / l1)) # 1826 倍 (2)像素 vs 感知(pixel_vs_perceptual.py):多尺度高斯差分特征加逐通道归一化,复现 LPIPS「在特征空间量距离」的结构。再次强调这是教学代理:本体用的是在 BAPPS 上学过权重的 VGG(见第 05 节真实代码)。 (3)GAN 损失与 warm-up(gan_schedule.py): def hinge_d(real, fake): return 0.5 * (np.mean(np.maximum(0.0, 1.0 - real)) + np.mean(np.maximum(0.0, 1.0 + fake))) def adopt_weight(weight, step, threshold=0, value=0.0): return value if step < threshold else weight (4)自适应权重(adaptive_weight.py):对迷你线性解码器用解析梯度(并用中心差分校验)算范数比: d_weight = np.clip(np.linalg.norm(g_rec) / (np.linalg.norm(g_gan) + 1e-4), 0.0, 1e4) print("早期 %.3f -> 后期 %.4f" % (w_early, w_late)) # 0.095 -> 0.003 (5)时间一致性(temporal_consistency.py):比较单帧空间 L1 与相邻帧差分误差,证明只有后者能抓到闪烁。 这些脚本的真实输出已散落在第 03 节,文末附录给出可直接运行的完整源码。 05. 工业级实现对照 最小实现是为了讲清原理,生产代码有几处关键的工程化。最权威的参照是 VQGAN 的损失模块(知识树锚点): CompVis/taming-transformers/taming/modules/losses/vqperceptual.py → VQLPIPSWithDiscriminator.forward(以 2026-09 的实现为准) 对照本文的四项,它的生成器一步几乎是公式的逐行翻译: rec_loss = torch.abs(inputs - reconstructions) # L1 像素 if self.perceptual_weight > 0: p_loss = self.perceptual_loss(inputs, reconstructions) # LPIPS rec_loss = rec_loss + self.perceptual_weight * p_loss nll_loss = torch.mean(rec_loss) logits_fake = self.discriminator(reconstructions) g_loss = -torch.mean(logits_fake) # hinge 生成损失 d_weight = self.calculate_adaptive_weight(nll_loss, g_loss, last_layer) loss = nll_loss + d_weight * disc_factor * g_loss \ + self.codebook_weight * codebook_loss.mean() # VQ:codebook;连续 VAE 换成 KL 与最小实现的差异,每一处都有来由: 重建与感知合并成 nll_loss 一起算梯度:自适应权重需要对「合并后的重建项」和「GAN 项」分别在最后一层求梯度(torch.autograd.grad(..., retain_graph=True)),所以 LPIPS 不是独立加权、而是并进重建项。 双优化器、两次前向:optimizer_idx==0 更新生成器(含重建、LPIPS、codebook、GAN),==1 才更新判别器;判别器那一路对输入 .detach(),不让梯度流回解码器。 LPIPS 本体(taming/modules/losses/lpips.py):先过一个 ScalingLayer(把通常约定为 $[-1,1]$ 的输入变到 VGG 的归一化统计,shift/scale 是写死的常数),再取 VGG16 的 5 个 relu 特征,逐通道归一化、过学来的 $1\times1$ 卷积、空间平均、跨层求和 —— 正是公式 3.3 的完整实现,权重从 BAPPS 预训练 ckpt 加载且全程冻结。 连续视频 VAE 用 KL 约束替代离散码本损失,但具体实现还可能改变重建项归约、学习似然尺度、判别器和时序结构,不能机械替换一行就认为目标完全相同。 当代视频 VAE 的公开配方(权重都来自各自技术报告,不要跨项目照抄,口径不同): 模型 损失组合与权重(论文原文) HunyuanVideo (2412.03603) $\mathrm{L_1}+0.1\,\mathrm{L_{lpips}}+0.05\,\mathrm{L_{adv}}+10^{-6}\,\mathrm{L_{kl}}$;判别器做随机缩放加时间维扩展,视频与图像从零联合训练 LTX-Video (2501.00103) 像素 MSE 加 Video-DWT(小波域 L1)加 LPIPS 加 Reconstruction-GAN;并讨论 causal /non-causal VAE 的取舍 Seedance 1.0 (2506.09113) L1 加 KL 加 LPIPS 加对抗损失;用类 PatchGAN 的混合判别器同时约束外观与运动 H3AE (2504.10567) 反方证据:判别类损失收益小却显著拖慢训练,主张先用 L1+KL 收敛、再用潜空间一致性损失微调 表中只有 HunyuanVideo 给出了这里列出的明确系数,其他项目采用不同目标或归一化,不能据此称权重“高度趋同”。toy 的数值只解释口径为何重要,不能从 1826 倍损失比推导出真实训练必需的 $10^{-6}$ 权重。 06. 代价与边界 每一项都在解决一类失真,也都引入新的代价: L1/L2:稳、好训,但 L2 糊、L1 颗粒重;它们只能保证「像素对」,保证不了「看着真」。 KL:去掉它可能使后验更确定、尺度约束变弱,但没有 KL 的自编码器或 VQ latent 也可以训练扩散模型;但权重过大、瓶颈太紧,重建细节会被牺牲。$10^{-6}$ 这种小权重是「弱正则」,依赖特定归一化口径,换套实现可能就要重新定标。 LPIPS:贴人眼,但它的「审美」被冻结在 VGG 的自然图像特征里 —— 对医学影像、动画、线稿等域外数据可能偏置;它偏纹理,有时会鼓励「看起来有细节」的伪纹理。 GAN:是清晰度和真实感的主要来源,也几乎是所有训练不稳的来源:需要 warm-up、谱归一化、限制判别器更新次数、自适应权重;并且提升感知质量往往以 PSNR 下降为代价(感知 — 失真权衡,perception–distortion tradeoff),这不是没训好,而是规律。 视频时间项 / 3D 判别器:能压闪烁,但显著增算力、增训练时长;光流估计本身在遮挡和大运动处会出错,warp 损失可能误伤真实运动。 边界也要讲清:H3AE 等近期工作指出,在高压缩 VAE 上判别类损失的边际收益可能撑不起它的训练成本,先用重建加 KL、后期再针对性微调是更划算的路线。四项全开不是政治正确—— 数据域、压缩率、训练阶段不同,最优组合也不同,应当用消融实验决定。 07. 经典论文脉络 Auto-Encoding Variational Bayes(VAE,arXiv:1312.6114,2013):提出 ELBO、重参数化与连续高斯潜变量,KL 正则的源头。 Neural Discrete Representation Learning(VQ-VAE,arXiv:1711.00937,2017):改走离散码本,用 codebook/commitment 损失替代 KL,是 VQGAN 的前身。 Image-to-Image Translation with Conditional Adversarial Networks(pix2pix / PatchGAN,arXiv:1611.07004,2017):把判别器做成局部 patch 判定器,确立了「局部对抗项与像素重建互补」的分工。 The Unreasonable Effectiveness of Deep Features as a Perceptual Metric(LPIPS,arXiv:1801.03924,CVPR 2018):用 BAPPS 人类偏好数据证明深度特征距离远胜 PSNR/SSIM,并学出逐层权重。 Taming Transformers for High-Resolution Image Synthesis(VQGAN,arXiv:2012.09841,2021):L1、LPIPS、PatchGAN、codebook 四件套定型,配套自适应 GAN 权重与 warm-up,是本文工业对照的母本。 CogVideoX(arXiv:2408.06072)与 HunyuanVideo(arXiv:2412.03603):把这套组合搬到 3D 因果视频 VAE,后者给出明确的四项权重和时空判别器设计。 H3AE(arXiv:2504.10567,2025):对「判别损失是否值得」提出反方证据,代表这条线仍在演进。 08. 常见误解 「KL 权重抄 HunyuanVideo 的 $10^{-6}$ 就行」:错。权重取决于你的 KL 是 sum 还是 mean、latent 与像素各有多少元素、有没有做 loss balancing。本文实测同一组数据 sum 口径 KL 是 mean 口径 L1 的 1826 倍;换个压缩率或归一化,$10^{-6}$ 可能过大或过小。先统一口径,再谈数字。 「LPIPS 就是拿 VGG 特征算 L2」:漏了两个关键件 —— 沿通道的归一化和在 BAPPS 上学出来的 $1\times1$ 权重;输入还要先过 scaling layer 换到 VGG 的数值域。少了归一化,距离会被少数高响应通道主导。 「生成器 GAN 损失算出来是负数,训练崩了」:hinge 生成器目标下 $\mathcal{L}_G=-\mathbb{E}[D(r)]$(logistic non-saturating 常写成 $\mathbb E[\mathrm{softplus}(-D(r))]$,是另一种公式),为负恰恰说明重建帧已经把判别器骗到打正分,是预期现象;该盯的是梯度和平衡,不是损失正负。 「判别器越强,重建越清晰」:hinge margin 饱和会让判别器损失梯度为 0,但生成器目标对 fake logit 的导数仍为 −1,不能把 D loss=0 当成生成器梯度消失的证据。正确姿势是晚启动、谱归一化、限制判别器更新次数、用自适应权重。 「L1/L2 越低画面越好」:L2 的最优解是条件均值,越低往往越糊;清晰度是用 GAN / 感知项换来的,并伴随 PSNR 下降。要同时看像素指标和感知 / 对抗指标。 「视频 VAE 把图像四项 loss 直接套到 3D 卷积上就行」:单帧损失测不出帧间闪烁(实测单帧 L1 几乎相同、时序误差差 393 倍)。应额外评估时序误差,按消融结果决定是否需要时空判别器、Video-DWT 或显式帧间约束;时间网络结构本身也能学习一致性。 「四项应该从头一起训」:主流做法是重建加 KL(有时加 LPIPS)先预热,第 disc_start 步才开 GAN;H3AE 甚至主张慎用判别损失。晚启动是一种稳定训练的办法,是否必需以及何时启动需依据具体实验。 09. 动手验证 5 个脚本都在本文附录,仅需 numpy(pip install numpy);第 6 个 make_figures.py 额外需要 matplotlib,用来重画本文配图。预期关键结果如下(随机种子已固定): 运行 vaeloss_terms.py:看到 L1 约 0.0316、KL (sum) 约 57.7、KL (mean) 约 0.225,以及「KL (sum)/L1 约 1826 倍」和 HunyuanVideo 权重配方下各项贡献 —— 建立「权重是在配平量纲」的直觉。 运行 pixel_vs_perceptual.py:亮度偏移与同能量噪声的像素 MSE 几乎相等(差小于 1%),但感知距离相差约三个数量级(噪声约为偏移的 $8\times10^3$ 倍)。 运行 gan_schedule.py:warm-up 在 step 2000 起跳;判别器「过强」时 D (hinge)=0.0000 而 G 约 5.97,同时检查脚本新增的 logit 导数:D margin 梯度归零,G 对 fake logit 的导数仍非零。 运行 adaptive_weight.py:先看有限差分与解析梯度完全一致,再看 d_weight 从训练早期 0.095 降到近收敛 0.003(约 32 倍)。 运行 temporal_consistency.py:闪烁与稳定两种重建的单帧 L1 约 0.063 几乎相同,时序 MSE 却差约 393 倍 —— 理解视频为何要单独约束时间维。 运行 make_figures.py(这个需要额外装 pip install matplotlib):重新生成本文三张配图 —— 量级账本(加权前 vs 加权后)、闪烁轨迹与指标对比、自适应权重随收敛下降的曲线。把 vaeloss_terms.py 里的 KL 权重从 $10^{-6}$ 调大两个数量级再看图 1,KL 项由约 $5.8\times10^{-5}$ 升到 $5.8\times10^{-3}$,仍低于这里约 0.0316 的 L1,不能声称已经压过所有项。 建议进一步动手:把 vaeloss_terms.py 里的 KL 从 sum() 改成 mean(),观察 HunyuanVideo 配方里 KL 项的相对权重需要相应放大多少倍,就能切身体会「权重不能跨口径照抄」。 10. 延伸阅读 读这篇之前建议先看已发布的前置: 变分下界与重参数化:KL 解析式与重参数化技巧的完整推导。 VAE 结构与训练目标:编码器 / 解码器、scaling factor、后验坍缩。 视频 VAE 的时空压缩结构:3D 因果卷积、时间 / 空间压缩比与分块解码。 读完这篇可以继续看: 视频生成评测:VBench 与人工验收(还没写):学会在分维度指标上看出「总分没变、时序一致性掉了」这类压缩副作用。 FID / CLIP Score 到底测了什么:感知与分布距离指标的适用边界,与本文 LPIPS/GAN 的质量观互补。 离散化表征:VQ-VAE 与 VQGAN(还没写):本文连续 KL-VAE 的「离散码本」对照版本,codebook 损失替代 KL。 附录:完整代码 09 节用到的脚本全文如下(vaeloss_terms.py、make_figures.py、pixel_vs_perceptual.py、gan_schedule.py、adaptive_weight.py、temporal_consistency.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 vaeloss_terms.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """视频 VAE 四类损失的最小可运行对照:L1 / KL / 感知 / GAN。 只依赖 numpy,直接 `python vaeloss_terms.py` 复现。 这一步不训练任何网络,只把四项损失放到同一个合成视频张量上算一遍, 回答一个最容易把人带沟里的问题:它们的「数值量级」根本不在一个频道上, 为什么工业配方里 KL 的权重会小到 1e-6。 """ import numpy as np rng = np.random.default_rng(0) # ---------------------------------------------------------------------- # 一个迷你的多尺度「特征提取器」,用来演示感知损失的*结构*。 # 注意:它不是 LPIPS 本体——LPIPS 用的是在 BAPPS 上学过权重的 VGG/Alex # 特征(见工业对照一节的 taming/lpips.py)。这里用可分离高斯 + 差分(DoG) # 构造确定性的多尺度边缘特征,目的是让「在特征空间而非像素空间比较」 # 这件事可以离线、免权重地跑起来。 # ---------------------------------------------------------------------- def _gauss_kernel(size=5, sigma=1.0): ax = np.arange(size) - (size - 1) / 2.0 k = np.exp(-(ax ** 2) / (2 * sigma ** 2)) return k / k.sum() def _conv_separable(img, k): """对最后两轴(H, W)做可分离卷积;img 形状 [..., H, W],边界 reflect。""" pad = len(k) // 2 x = np.pad(img, ((0, 0),) * (img.ndim - 2) + ((pad, pad), (0, 0)), mode="reflect") acc = np.zeros_like(img) for i, w in enumerate(k): acc += w * x[..., i:i + img.shape[-2], :] x = np.pad(acc, ((0, 0),) * (img.ndim - 2) + ((0, 0), (pad, pad)), mode="reflect") acc = np.zeros_like(img) for j, w in enumerate(k): acc += w * x[..., :, j:j + img.shape[-1]] return acc def _normalize_channels(feat, eps=1e-10): # 对应 LPIPS 的 normalize_tensor:沿通道维归一化 norm = np.sqrt(np.sum(feat ** 2, axis=1, keepdims=True)) return feat / (norm + eps) def feature_stack(x): """x: [B, C, T, H, W] -> 多尺度边缘特征列表(先把 T 折叠进 batch)。""" B, C, T, H, W = x.shape f = x.transpose(0, 2, 1, 3, 4).reshape(B * T, C, H, W) k = _gauss_kernel(5, 1.0) b1 = _conv_separable(f, k) b2 = _conv_separable(_conv_separable(b1, k), k) dog1 = f - b1 # 高频细节 dog2 = b1 - b2 # 中频边缘 feats = [_normalize_channels(f), _normalize_channels(dog1), _normalize_channels(dog2)] return feats def perceptual_proxy(x, y): """结构对齐 LPIPS:逐层归一化特征差的平方,空间平均后跨层求和。""" total = 0.0 for fx, fy in zip(feature_stack(x), feature_stack(y)): total += np.mean((fx - fy) ** 2) return total def make_video(B=2, T=8, H=32, W=32): """合成一段有运动内容的视频:移动的亮圆 + 网格背景,取值[0,1]。""" x = np.zeros((B, 3, T, H, W), dtype=np.float64) yy, xx = np.mgrid[0:H, 0:W] for b in range(B): for t in range(T): cx = W * (0.3 + 0.5 * t / (T - 1)) cy = H * (0.3 + 0.15 * np.sin(2 * np.pi * t / T)) circle = ((xx - cx) ** 2 + (yy - cy) ** 2) < (H * 0.12) ** 2 grid = ((xx // 8 + yy // 8) % 2) * 0.15 frame = 0.2 + grid frame[circle] = 0.95 x[b, 0, t] = frame x[b, 1, t] = frame * 0.9 x[b, 2, t] = frame * 0.7 return x def degrade(x): """重建结果:轻微模糊 + 小噪声 + 偏色,模拟一个训练中段的解码器。""" B, C, T, H, W = x.shape flat = x.transpose(0, 2, 1, 3, 4).reshape(B * T, C, H, W) k = _gauss_kernel(3, 0.8) out = _conv_separable(flat, k) out = out + rng.normal(0, 0.02, out.shape) out[:, 0] += 0.03 # 轻微红色偏置 return np.clip(out, 0, 1).reshape(B, T, C, H, W).transpose(0, 2, 1, 3, 4) def main(): x = make_video() r = degrade(x) print("视频张量 x / r :", x.shape, " 取值范围 %.2f~%.2f" % (x.min(), x.max())) # 编码器输出的后验 q(z|x):latent 在时间压缩4x、空间压缩8x B, C, T, H, W = x.shape Cz, Tz, Hz, Wz = 4, T // 4, H // 8, W // 8 mu = 0.3 * rng.standard_normal((B, Cz, Tz, Hz, Wz)) logvar = -1.0 + 0.2 * rng.standard_normal((B, Cz, Tz, Hz, Wz)) n_pix = B * C * T * H * W n_lat = mu.size print("latent 形状 :", mu.shape, " 像素数=%d latent数=%d\n" % (n_pix, n_lat)) # ① 像素重建 L1(对全部元素求平均,量纲与像素一致,O(0.01~0.1)) l1 = np.mean(np.abs(x - r)) # ② KL 解析式(标准正态先验)。注意它天然是「求和」式 kl_per_element = -0.5 * (1 + logvar - mu ** 2 - np.exp(logvar)) kl_sum = kl_per_element.sum() kl_mean = kl_per_element.mean() # ③ 感知损失(多尺度特征距离,量级被归一化压在 O(0.01)) l_p = perceptual_proxy(x, r) # ④ GAN 生成器损失。这里直接喂一组判别器对假图的打分 logits # hinge 生成损失 = -mean(logits_fake);先假设判别器刚起步、打分偏正 logits_fake = rng.normal(0.3, 0.5, size=64) g_loss = -np.mean(logits_fake) print("%-28s %10s" % ("损失项", "原始数值")) print("-" * 40) print("%-28s %10.5f" % ("L1 重建 (mean)", l1)) print("%-28s %10.5f" % ("KL (sum,常见写法)", kl_sum)) print("%-28s %10.5f" % ("KL (mean,换口径)", kl_mean)) print("%-28s %10.5f" % ("感知 proxy", l_p)) print("%-28s %10.5f" % ("GAN g=-mean(D(r))", g_loss)) print("\n--- 为什么 KL 权重能小到 1e-6 ---") # 若直接把 sum 口径的 KL 与 mean 口径的 L1 相加,KL 会凭元素数量碾压: print("KL(sum)/L1 的量级倍数 : %.1fx" % (kl_sum / l1)) print("HunyuanVideo 配方 L1 + 0.1*LPIPS + 0.05*GAN + 1e-6*KL :") total = l1 + 0.1 * l_p + 0.05 * g_loss + 1e-6 * kl_sum print(" 各项贡献: L1=%.5f LPIPS=%.5f GAN=%.5f KL=%.6f" % (l1, 0.1 * l_p, 0.05 * g_loss, 1e-6 * kl_sum)) print(" 合计 = %.5f" % total) print("\n结论:权重不是玄学,是在给「不同口径、不同元素数、不同量纲」的项找平。") if __name__ == "__main__": main() make_figures.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """生成「视频 VAE 的常见 loss 组合」的三张解释图。 数值一律从同目录的三个脚本里取(vaeloss_terms / temporal_consistency / adaptive_weight),这里只负责画——改了那边这里要重跑,免得图和正文数字打架。 三张图分别回答: 1. 四项损失的原始量级差着几个数量级,工业配方加权后为什么能凑到一起 2. 单帧指标完全分不开的两种重建,时间维损失一眼看穿 3. 自适应 GAN 权重为什么随训练自动变小 只依赖 numpy + matplotlib。跑法:python make_figures.py """ import textwrap from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np import adaptive_weight as AW import temporal_consistency as TC import vaeloss_terms as VT ROOT = Path(__file__).resolve().parents[1] OUT = ROOT / "figures" OUT.mkdir(parents=True, exist_ok=True) plt.rcParams.update({ "font.sans-serif": ["PingFang SC", "Arial Unicode MS", "DejaVu Sans"], "axes.unicode_minus": False, "figure.dpi": 160, "savefig.bbox": "tight", }) INK = "#1f2937" MUTE = "#6b7280" C_PIX = "#2f6fb0" # 像素项:蓝 C_PERC = "#8b5cf6" # 感知项:紫 C_GAN = "#e0a03c" # 对抗项:橙 C_KL = "#d1495b" # KL:红 C_A = "#d1495b" # 闪烁重建 C_B = "#2f9e6f" # 稳定重建 C_GT = "#1f2937" # 真值 def style(ax, title, xlabel=None, ylabel=None): ax.set_title(title, fontsize=11.5, color=INK, pad=10, loc="left") if xlabel: ax.set_xlabel(xlabel, fontsize=10, color=MUTE) if ylabel: ax.set_ylabel(ylabel, fontsize=10, color=MUTE) ax.tick_params(colors=MUTE, labelsize=9) for s in ("top", "right"): ax.spines[s].set_visible(False) for s in ("left", "bottom"): ax.spines[s].set_color("#d1d5db") ax.grid(alpha=0.25, linewidth=0.6, axis="y") ax.set_axisbelow(True) def footer(fig, text, width=118): """把「这张图要看什么」放到坐标轴下方。 必须放在 y<0 的位置:bbox_inches="tight" 会把负坐标的 artist 一起收进来, 放在 0~0.05 之间的话会和 x 轴标签叠在一起。 """ wrapped = "\n".join(textwrap.wrap(text, width=width)) fig.text(0.012, -0.13, wrapped, fontsize=8.5, color=MUTE, va="top", ha="left", linespacing=1.6) # ────────────────────────────────────────────────────────────────────── # 图 1:四项损失的量级账本 # ────────────────────────────────────────────────────────────────────── def fig_magnitudes(): x, r = VT.make_video(), VT.degrade(VT.make_video()) # 与 vaeloss_terms.main 完全同口径重算一遍(rng 序列一致) B, C, T, H, W = x.shape Cz, Tz, Hz, Wz = 4, T // 4, H // 8, W // 8 mu = 0.3 * VT.rng.standard_normal((B, Cz, Tz, Hz, Wz)) logvar = -1.0 + 0.2 * VT.rng.standard_normal((B, Cz, Tz, Hz, Wz)) l1 = float(np.mean(np.abs(x - r))) kl_el = -0.5 * (1 + logvar - mu ** 2 - np.exp(logvar)) kl_sum, kl_mean = float(kl_el.sum()), float(kl_el.mean()) lp = float(VT.perceptual_proxy(x, r)) gan = abs(float(-np.mean(VT.rng.normal(0.3, 0.5, size=64)))) fig, axes = plt.subplots(1, 2, figsize=(11.6, 4.4)) ax = axes[0] names = ["L1 重建", "感知 proxy", "KL(mean)", "KL(sum)", "|GAN|"] vals = [l1, lp, kl_mean, kl_sum, gan] cols = [C_PIX, C_PERC, C_KL, C_KL, C_GAN] bars = ax.bar(names, vals, color=cols, width=0.62) ax.set_yscale("log") ax.set_ylim(1e-3, 3e2) for b, v in zip(bars, vals): ax.text(b.get_x() + b.get_width() / 2, v * 1.25, f"{v:.4g}", ha="center", fontsize=8.5, color=INK) style(ax, "加权前:同一视频上四项损失的原始数值", None, "损失值(对数轴)") ax = axes[1] wnames = ["1.0 × L1", "0.1 × 感知", "0.05 × GAN", "1e-6 × KL(sum)"] wvals = [l1, 0.1 * lp, 0.05 * gan, 1e-6 * kl_sum] wcols = [C_PIX, C_PERC, C_GAN, C_KL] bars = ax.bar(wnames, wvals, color=wcols, width=0.62) ax.set_yscale("log") ax.set_ylim(1e-6, 1e-1) for b, v in zip(bars, wvals): ax.text(b.get_x() + b.get_width() / 2, v * 1.6, f"{v:.4g}", ha="center", fontsize=8.5, color=INK) ax.axhline(l1, color=MUTE, ls=":", lw=1.0) ax.text(2.6, l1 * 1.5, "L1 的量级", fontsize=8, color=MUTE) style(ax, "加权后:HunyuanVideo 配方下各项的真实贡献", None, "对总损失的贡献(对数轴)") fig.suptitle("图 1 损失归约口径与加权贡献:标量大小不等于梯度大小", fontsize=12, color=INK, x=0.012, ha="left", y=1.02) footer(fig, "要看什么:左图是四项损失在同一个合成视频上的原始数值——KL 按教科书口径对 latent 维" "求和,一上来就是 L1 的 1826 倍;右图是套上工业权重之后各项的真实贡献," "贡献仍不同量级(KL 约5.8e-5,而L1约0.0316);这些权重不能由toy损失比唯一推导。" "KL 权重小到 1e-6 不是不重要,是在补偿「求和口径 vs 平均口径」的元素数之差。") fig.savefig(OUT / "loss_magnitudes.png") plt.close(fig) print(f"[图1] L1={l1:.5f} KL_sum={kl_sum:.5f} 倍数={kl_sum / l1:.0f}x;" f"加权后贡献 L1={l1:.5f} KL={1e-6 * kl_sum:.2e}") # ────────────────────────────────────────────────────────────────────── # 图 2:单帧指标看不见的闪烁 # ────────────────────────────────────────────────────────────────────── def fig_temporal(): x = TC.make_sequence() # [T,1,H,W] fixed = TC.SIGMA * TC.rng.standard_normal(x.shape[1:])[None] rA = x + TC.SIGMA * TC.rng.standard_normal(x.shape) rB = x + fixed + 0.05 * TC.SIGMA * TC.rng.standard_normal(x.shape) # 取一条穿过运动边缘的水平线上的一个像素,看它随帧的取值 t_axis = np.arange(TC.T) y0, x0 = TC.H // 2, 30 gt = x[:, 0, y0, x0] a = rA[:, 0, y0, x0] b = rB[:, 0, y0, x0] fig, axes = plt.subplots(1, 2, figsize=(11.6, 4.3)) ax = axes[0] ax.plot(t_axis, gt, lw=2.6, color=C_GT, label="真值", zorder=3) ax.plot(t_axis, a, lw=1.2, color=C_A, alpha=0.9, label="A 逐帧独立噪声(闪)") ax.plot(t_axis, b, lw=1.2, color=C_B, alpha=0.9, label="B 跨帧恒定退化(稳)") ax.set_ylim(gt.min() - 4 * TC.SIGMA, gt.max() + 4 * TC.SIGMA) ax.legend(fontsize=8.5, frameon=False, loc="upper left") style(ax, "同一个像素随帧的变化:A 在真值附近抖,B 整体平移但不抖", "帧 t", "像素值") ax = axes[1] mets = ["单帧空间 L1", "时序差分 L1", "时序差分 MSE"] va = [TC.spatial_l1(rA, x), TC.temporal_error(rA, x), TC.temporal_mse(rA, x)] vb = [TC.spatial_l1(rB, x), TC.temporal_error(rB, x), TC.temporal_mse(rB, x)] xg = np.arange(len(mets)) w = 0.36 ba = ax.bar(xg - w / 2, va, w, color=C_A, label="A 闪") bb = ax.bar(xg + w / 2, vb, w, color=C_B, label="B 稳") ax.set_yscale("log") ax.set_ylim(1e-5, 1) for bars in (ba, bb): for b_ in bars: ax.text(b_.get_x() + b_.get_width() / 2, b_.get_height() * 1.5, f"{b_.get_height():.4g}", ha="center", fontsize=8, color=INK) ax.set_xticks(xg) ax.set_xticklabels(mets) ax.legend(fontsize=8.5, frameon=False, loc="upper left") style(ax, "三种指标下 A、B 的差距:第一列分不开,第三列差 393 倍", None, "误差(对数轴)") fig.suptitle("图 2 单帧 L1 完全分不开的两种重建,时序差分 MSE 差了 393 倍", fontsize=12, color=INK, x=0.012, ha="left", y=1.02) footer(fig, "要看什么:左图是同一个像素在 12 帧上的取值——A(红)每一帧都在真值附近独立抖动," "连起来播放就是闪烁;B(绿)带一个恒定偏移但帧间几乎不动。" "右图是量化结论:单帧空间 L1 下 A=0.0639、B=0.0634,指标根本分不出谁好谁坏;" "换时序差分 MSE,A 是 B 的 393 倍。只看单帧指标,闪烁是隐形的。") fig.savefig(OUT / "temporal_flicker.png") plt.close(fig) print(f"[图2] 单帧L1 A={va[0]:.5f}/B={vb[0]:.5f};时序MSE A={va[2]:.5f}/B={vb[2]:.5f}" f"({va[2] / max(vb[2], 1e-12):.0f}x)") # ────────────────────────────────────────────────────────────────────── # 图 3:自适应 GAN 权重随训练自动变小 # ────────────────────────────────────────────────────────────────────── def fig_adaptive(): b = np.zeros(AW.D) # 先按 adaptive_weight.main 的抽随机顺序取两个阶段点,保证与正文数字一致 W_early = AW.rng.standard_normal((AW.dz, AW.D)) * 0.05 W_late = AW.W_star + AW.rng.standard_normal((AW.dz, AW.D)) * 0.01 g_rec, g_gan = AW.grads(W_early, b) w_early = AW.adaptive_weight(g_rec, g_gan) g_rec2, g_gan2 = AW.grads(W_late, b) w_late = AW.adaptive_weight(g_rec2, g_gan2) # 扫描曲线用独立的 rng,不扰动上面的阶段点 sweep_rng = np.random.default_rng(7) scales = np.logspace(-4, -0.3, 14) weights = [] for s in scales: W = AW.W_star + s * sweep_rng.standard_normal((AW.dz, AW.D)) g_rec, g_gan = AW.grads(W, b) weights.append(AW.adaptive_weight(g_rec, g_gan)) weights = np.array(weights) fig, ax = plt.subplots(figsize=(7.8, 4.4)) ax.plot(scales, weights, marker="o", ms=4.5, lw=2.0, color=C_GAN, label=r"自适应权重 $d_{\mathrm{weight}}$") ax.axvline(0.05, color=MUTE, ls=":", lw=1.0) ax.text(0.055, w_early * 2.2, "训练早期\n(初始化附近)", fontsize=8.5, color=MUTE) ax.axvline(0.01, color=MUTE, ls=":", lw=1.0) ax.text(0.0105, w_late * 0.06, "训练后期\n(近收敛)", fontsize=8.5, color=MUTE) for s, w_ in [(0.05, w_early), (0.01, w_late)]: ax.plot([s], [w_], marker="s", ms=7, color=C_PIX, zorder=5) ax.annotate(f"{w_early:.3f}", xy=(0.05, w_early), xytext=(0.075, w_early * 1.6), fontsize=9, color=C_PIX) ax.annotate(f"{w_late:.4f}", xy=(0.01, w_late), xytext=(0.0125, w_late * 0.4), fontsize=9, color=C_PIX) ax.set_xscale("log") ax.set_yscale("log") ax.legend(fontsize=9, frameon=False, loc="upper right") style(ax, "重建越接近收敛,GAN 项被自动调得越小", "解码器离最优解的距离(权重扰动幅度,对数轴)", r"自适应权重 $d_{\mathrm{weight}}$(对数轴)") fig.suptitle("图 3 固定 λ 会两头翻车:早期压不住重建、后期压不住 GAN", fontsize=12, color=INK, x=0.012, ha="left", y=1.03) footer(fig, "要看什么:横轴是解码器离最优解有多远(越靠左越接近收敛),纵轴是 VQGAN 的" "自适应权重 ||∇L_rec|| / ||∇L_GAN||。它随残差缩小近似线性下降——重建收敛后" f"从 {w_early:.3f} 掉到 {w_late:.4f}(约 {w_early / max(w_late, 1e-12):.0f} 倍),GAN 项被同步调小。" "若用固定权重,早期 GAN 抢不过巨大的重建梯度,后期重建梯度变小、GAN 又反过来" "压过重建——这条斜线就是在消除这个漂移。") fig.savefig(OUT / "adaptive_weight_curve.png") plt.close(fig) print(f"[图3] d_weight: 早期 {w_early:.4f} -> 后期 {w_late:.4f}" f"({w_early / max(w_late, 1e-12):.0f}x)") def main(): fig_magnitudes() fig_temporal() fig_adaptive() print(f"[OK] 三张图已写入 {OUT}") for f in sorted(OUT.glob("*.png")): print(f" {f.name} {f.stat().st_size / 1024:.0f} KB") if __name__ == "__main__": main() pixel_vs_perceptual.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """复现 LPIPS 论文最核心的观察:相同的像素误差,感知质量可以天差地别。 只依赖 numpy,直接 `python pixel_vs_perceptual.py` 复现。 我们构造三种「重建」,让其中两种的像素 MSE 完全相等: A. 全局亮度偏移 —— 人眼对缓慢的整体明暗变化相当不敏感; B. 同等 L2 能量的逐点噪声 —— 直接污染纹理与细节,观感很差; C. 高斯模糊 —— 抹掉高频细节。 然后分别在「像素空间」和一个多尺度特征空间里比较距离。 """ import numpy as np rng = np.random.default_rng(1) def _gauss_kernel(size=5, sigma=1.0): ax = np.arange(size) - (size - 1) / 2.0 k = np.exp(-(ax ** 2) / (2 * sigma ** 2)) return k / k.sum() def _conv_separable(img, k): pad = len(k) // 2 x = np.pad(img, ((0, 0),) * (img.ndim - 2) + ((pad, pad), (0, 0)), mode="reflect") acc = np.zeros_like(img) for i, w in enumerate(k): acc += w * x[..., i:i + img.shape[-2], :] x = np.pad(acc, ((0, 0),) * (img.ndim - 2) + ((0, 0), (pad, pad)), mode="reflect") acc = np.zeros_like(img) for j, w in enumerate(k): acc += w * x[..., :, j:j + img.shape[-1]] return acc def _norm(feat, eps=1e-10): return feat / (np.sqrt(np.sum(feat ** 2, axis=1, keepdims=True)) + eps) def perceptual_proxy(x, y): """多尺度高通特征上的归一化距离(LPIPS 的结构代理,非本体)。""" def feats(z): k = _gauss_kernel(5, 1.0) b1 = _conv_separable(z, k) b2 = _conv_separable(_conv_separable(b1, k), k) return [_norm(z), _norm(z - b1), _norm(b1 - b2)] return sum(np.mean((a - b) ** 2) for a, b in zip(feats(x), feats(y))) def make_textured_image(H=64, W=64): """一张同时含锐边、细密纹理和平滑区域的图。""" yy, xx = np.mgrid[0:H, 0:W].astype(np.float64) img = 0.5 + 0.25 * np.sin(xx / 2.0) * np.sin(yy / 2.0) # 细密纹理 img += 0.3 * (xx > W / 2) # 一条锐边 checker = ((xx // 4 + yy // 4) % 2) * 0.15 # 棋盘格 img += checker img = np.clip(img, 0, 1) return np.stack([img, img * 0.92, img * 0.8], axis=0)[None] # [1,C,H,W] def mse(a, b): return float(np.mean((a - b) ** 2)) def main(): x = make_textured_image() delta = 0.12 # 统一的 L2 误差能量 rA = x + delta # A:全局亮度偏移(不 clip,保证误差严格等于 delta) rB = x + rng.normal(0, delta, x.shape) # B:零均值逐点噪声,标准差=delta rC = _conv_separable(x, _gauss_kernel(7, 1.6)) # C:模糊 rows = [ ("A 亮度偏移(整体+%.2f)" % delta, rA), ("B 逐点噪声(σ=%.2f)" % delta, rB), ("C 高斯模糊", rC), ] print("%-26s %12s %12s %14s" % ("重建方式", "像素MSE", "感知距离", "感知/像素 比")) print("-" * 68) for name, r in rows: pm = mse(x, r) pp = perceptual_proxy(x, r) print("%-26s %12.5f %12.5f %14.2f" % (name, pm, pp, pp / (pm + 1e-12))) mseA, mseB = mse(x, rA), mse(x, rB) pA, pB = perceptual_proxy(x, rA), perceptual_proxy(x, rB) print("\n关键对照:A 与 B 的像素 MSE 相差仅 %.2f%%," % (100 * abs(mseA - mseB) / mseA)) print("但特征空间里 B(噪声)的感知距离是 A(亮度偏移)的 %.1f 倍。" % (pB / pA)) print("高通/多尺度特征天然滤掉直流亮度、放大纹理污染——这正是 LPIPS 比 MSE 更贴人眼的原因。") if __name__ == "__main__": main() gan_schedule.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """GAN 两项的最小演示:hinge / vanilla 损失,以及判别器的 warm-up 调度。 只依赖 numpy,直接 `python gan_schedule.py` 复现。 视频 VAE 里对抗损失通常不是从头开:先用 L1+KL+LPIPS 把解码器训到大致成形, 第 disc_start 步才把判别器权重从 0 抬起来(taming 里的 adopt_weight)。 本脚本同时给出两种常见判别损失的数值,并展示「判别器过强」时的失衡信号。 """ import numpy as np rng = np.random.default_rng(3) def adopt_weight(weight, global_step, threshold=0, value=0.0): """taming-transformers 里的原逻辑:threshold 之前强制为 value(通常是0)。""" return value if global_step < threshold else weight def hinge_d(real, fake): return 0.5 * (np.mean(np.maximum(0.0, 1.0 - real)) + np.mean(np.maximum(0.0, 1.0 + fake))) def hinge_g(fake): return float(-np.mean(fake)) def vanilla_d(real, fake): softplus = lambda z: np.log1p(np.exp(-np.abs(z))) + np.maximum(z, 0.0) return 0.5 * (np.mean(softplus(-real)) + np.mean(softplus(fake))) def evaluate(real, fake, tag): print("%-28s D(hinge)=%.4f D(vanilla)=%.4f G(hinge)=%.4f 均值打分 real=%.2f fake=%.2f" % (tag, hinge_d(real, fake), vanilla_d(real, fake), hinge_g(fake), real.mean(), fake.mean())) def main(): # ① warm-up 调度:disc_start=2000 print("== adopt_weight 调度(disc_start=2000, 目标权重 0.05) ==") for step in (0, 1999, 2000, 5000): w = adopt_weight(0.05, step, threshold=2000) print(" step=%5d -> 对抗权重 = %.3f" % (step, w)) # ② 判别器在三种强弱下的损失 print("\n== 判别器强弱与损失信号 ==") # 平衡初期:真假打分都在 0 附近 real0 = rng.normal(0.0, 0.3, 256) fake0 = rng.normal(0.0, 0.3, 256) evaluate(real0, fake0, "起步:判别器分不清") # 健康:真≈+1,假≈-1,仍留梯度 real1 = rng.normal(1.2, 0.4, 256) fake1 = rng.normal(-1.2, 0.4, 256) evaluate(real1, fake1, "健康:适度拉开") # 过强:真≈+6,假≈-6,判别器 hinge 梯度为 0,但生成器对 fake logit 仍有梯度 real2 = rng.normal(6.0, 0.5, 256) fake2 = rng.normal(-6.0, 0.5, 256) evaluate(real2, fake2, "过强:margin 已饱和") d_fake_grad = 0.5 * (fake2 > -1.0) / fake2.size g_fake_grad = -np.ones_like(fake2) / fake2.size print("\nD 对 fake logit 梯度范数 = %.6f" % np.linalg.norm(d_fake_grad)) print("G 对 fake logit 梯度范数 = %.6f" % np.linalg.norm(g_fake_grad)) print("D margin 饱和不表示 G 梯度消失;参数梯度还取决于判别器输入雅可比。") if __name__ == "__main__": main() adaptive_weight.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """复现 VQGAN 的自适应 GAN 权重 calculate_adaptive_weight。 只依赖 numpy,直接 `python adaptive_weight.py` 复现。 固定 GAN 权重的痛点:训练早期解码器很烂,重建梯度很大;训练后期重建收敛、 梯度变小。若 λ 固定,GAN 梯度在后期会相对越来越强,甚至把重建带偏。 VQGAN 的做法是让两项在「最后一层」上的梯度范数之比来决定权重: d_weight = ||∇_last L_rec|| / (||∇_last L_gan|| + 1e-4),再 clamp 到 [0,1e4] 我们用一个可解析求导的迷你线性解码器,精确算出两个梯度范数, 并用有限差分抽查一个分量,证明数值不是凑出来的。 """ import numpy as np rng = np.random.default_rng(2) N, D, dz = 16, 24, 8 # 样本数、像素维、latent 维 Z = rng.standard_normal((N, dz)) W_star = rng.standard_normal((dz, D)) * 0.3 X = Z @ W_star # 真实数据:线性可完美拟合 # 一个固定的微型判别器打分函数 s(x)=x @ a(教学用,参数冻结) a = rng.standard_normal((D,)) def decode(W, b): return Z @ W + b def rec_loss(W, b): """重建项:L2(L1 同理,范数比机制不变)。""" return float(np.mean((decode(W, b) - X) ** 2)) def gan_g_loss(W, b): """生成器 hinge 损失 -mean(D(r))。""" s = decode(W, b) @ a return float(-np.mean(s)) def grads(W, b): R = decode(W, b) - X g_rec = (2.0 / (N * D)) * Z.T @ R # ∂L_rec/∂W(mean 对 N*D) # g=-mean_n(s_n),s_n=Σ_d xhat_nd·a_d:mean 只对 N,分母是 N,不是 N*D s_coef = -(1.0 / N) * Z.sum(axis=0) # ∂g/∂W_kd = -(1/N)(Σ_n z_nk) a_d g_gan = np.outer(s_coef, a) return g_rec, g_gan def adaptive_weight(g_rec, g_gan, cap=1e4): w = np.linalg.norm(g_rec) / (np.linalg.norm(g_gan) + 1e-4) return float(np.clip(w, 0.0, cap)) def finite_diff_check(W, b, eps=1e-6): """用中心差分抽查 W[0,0] 上的两个梯度,验证解析解。""" out = [] for fn in (rec_loss, gan_g_loss): Wp, Wm = W.copy(), W.copy() Wp[0, 0] += eps Wm[0, 0] -= eps out.append((fn(Wp, b) - fn(Wm, b)) / (2 * eps)) return out def stage(W, b, name): lrec = rec_loss(W, b) lgan = gan_g_loss(W, b) g_rec, g_gan = grads(W, b) w = adaptive_weight(g_rec, g_gan) print("%-22s L_rec=%8.5f L_gan=%8.4f ||∇rec||=%8.4f ||∇gan||=%7.4f d_weight=%7.4f" % (name, lrec, lgan, np.linalg.norm(g_rec), np.linalg.norm(g_gan), w)) return w def main(): b = np.zeros(D) # 阶段一:解码器刚初始化,离最优很远(重建残差大) W_early = rng.standard_normal((dz, D)) * 0.05 # 阶段二:接近收敛(在最优解上加很小扰动,残差小) W_late = W_star + rng.standard_normal((dz, D)) * 0.01 print("== 有限差分校验(W[0,0]) ==") fd_rec, fd_gan = finite_diff_check(W_early, b) g_rec, g_gan = grads(W_early, b) print("解析 ∂Lrec/∂W00=%.6f 差分=%.6f" % (g_rec[0, 0], fd_rec)) print("解析 ∂Lgan/∂W00=%.6f 差分=%.6f\n" % (g_gan[0, 0], fd_gan)) print("== 自适应权重随训练阶段变化 ==") w_early = stage(W_early, b, "训练早期(未收敛)") w_late = stage(W_late, b, "训练后期(近收敛)") print("\n重建收敛后 d_weight 从 %.3f 降到 %.4f(约 %.0f 倍),GAN 项被自动调小。" % (w_early, w_late, w_early / max(w_late, 1e-12))) print("若用固定 λ:早期 GAN 抢不过重建、后期 GAN 又可能压过重建——自适应权重就是在消除这个漂移。") if __name__ == "__main__": main() temporal_consistency.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """视频 VAE 为什么必须额外盯「时间维」:同样的单帧画质,闪烁程度可以天差地别。 只依赖 numpy,直接 `python temporal_consistency.py` 复现。 构造一段运动视频,再给两种重建: A. 逐帧独立噪声 —— 每一帧单独看误差不大,连起来播放却疯狂闪烁; B. 跨帧恒定的退化(同一固定纹理噪声)—— 单帧误差与 A 相同,但不闪。 单帧空间 L1 完全分不清 A、B;相邻帧差分的时序一致性损失一眼区分。 """ import numpy as np rng = np.random.default_rng(4) T, H, W = 12, 48, 64 SIGMA = 0.08 def make_sequence(): yy, xx = np.mgrid[0:H, 0:W].astype(np.float64) x = np.zeros((T, H, W)) for t in range(T): cx = 8 + 4 * t # 一条匀速移动的竖边 frame = 0.25 + 0.5 * (xx >= cx) frame += 0.1 * ((xx // 6 + yy // 6) % 2) x[t] = frame return np.clip(x, 0, 1)[:, None] # [T,1,H,W] def spatial_l1(r, x): return float(np.mean(np.abs(r - x))) def frame_diff(v): return v[1:] - v[:-1] def temporal_error(r, x): """相邻帧差分一致性:||Δr - Δx||(静止背景上 GT 帧差为0,闪烁直接显现)。""" return float(np.mean(np.abs(frame_diff(r) - frame_diff(x)))) def temporal_mse(r, x): return float(np.mean((frame_diff(r) - frame_diff(x)) ** 2)) def main(): x = make_sequence() fixed_noise = SIGMA * rng.standard_normal(x.shape[1:])[None] # [1,1,H,W] 跨帧恒定 rA = x + SIGMA * rng.standard_normal(x.shape) # 逐帧独立噪声 -> 闪烁 # 稳定退化:以跨帧恒定噪声为主,只掺 5% 的逐帧抖动(更贴近真实解码器) rB = x + fixed_noise + 0.05 * SIGMA * rng.standard_normal(x.shape) print("退化能量相同(σ=%.2f),比较单帧空间误差与时间维误差:\n" % SIGMA) print("%-22s %14s %16s %16s" % ("重建", "单帧空间L1", "时序差分L1", "时序差分MSE")) print("-" * 72) for name, r in (("A 逐帧独立噪声(闪)", rA), ("B 跨帧恒定退化(稳)", rB)): print("%-22s %14.5f %16.5f %16.5f" % (name, spatial_l1(r, x), temporal_error(r, x), temporal_mse(r, x))) print("\n单帧空间 L1 几乎相等(A=%.4f vs B=%.4f),但时序 MSE 上 A 约为 B 的 %.0f 倍。" % (spatial_l1(rA, x), spatial_l1(rB, x), temporal_mse(rA, x) / max(temporal_mse(rB, x), 1e-12))) print("这解释了视频 VAE 为何要在 2D 的 L1/LPIPS/GAN 之外:") print(" · 把判别器扩到时空(3D/PatchGAN、LTX 的 Video-DWT);") print(" · 或加相邻帧/warp 一致性损失,专门惩罚帧间抖动与闪烁。") if __name__ == "__main__": main()
2026年09月20日
5 阅读
0 评论
1 点赞
1
2
3
粤ICP备2021042327号