首页
应用
关于
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
累计撰写
203
篇文章
累计收到
8
条评论
首页
应用
栏目
AIGC
AIGC Daily Papers
AIGC Fundamentals
其他
职场经验复盘
多模态理解
购房/投资
阅读
广告
Segmentation
LeetCode
Pytorch
Python
Shell
C++
常用链接
页面
关于
搜索到
13
篇与
基础知识
的结果
2026-10-05
AIGC 基本功|扩散模型的跨步缓存复用-DiffCache
AIGC 基本功|扩散模型的跨步缓存复用-DiffCache 同样走完采样器的每一步,能不能少算几次昂贵的去噪网络?跨步缓存把这个问题拆成两件事:复用什么,以及什么时候必须刷新。本文用残差缓存的最小实验解释机制,再对照 DeepCache 和 TeaCache 的真实实现。实验是人工构造的数值系统,不是视频模型的质量或速度复现。 本文承接 DDIM 与高阶采样器、DiT 和 性能建模与 Profiling。默认读者知道采样器会反复调用网络,不要求先读缓存论文。论文版本与官方代码核对日期为 2026-10-05。 01. 为什么需要它 一次视频生成中,潜变量在逐步变化,文本条件通常固定,网络却要一遍遍穿过相同的模块。假如两次调用的深层特征几乎一样,重新计算这部分就可能浪费时间。但“几乎一样”只是观测,不能推出“任意一次都可以省”。运动切换、条件变化和采样末期的细节修整,都可能让陈旧特征的误差显现。 这里最具体的失败场景是固定间隔缓存:先完整算一次,之后连续复用,直到计数器要求刷新。如果变化突然出现在间隔中间,策略不会因为内容变了而提前刷新。反过来,如果某段变化很慢,固定间隔也会做多余的完整计算。问题在于日历式排期没有感知当前输入,周期短会损失收益,周期长会增加近似误差。 本文实验保留全部 40 次状态更新。精确版本执行 40 次完整残差计算;自适应阈值为 0.16 时,只执行 11 次,复用 29 次。终点相对误差为 0.008222。它证明了在这个特定数值系统中确实可以减少计算次数,同时也证明输出发生了变化。它不证明视频质量达标,更不证明真实 GPU 上能得到相同的提速。 缓存适合回答“这个步骤里,哪些计算可以暂时沿用”。减少采样步数则改变采样器的离散路径。两种方法可以组合,但组合后的误差要重新评估:单独用缓存通过了一组样例,不意味着换成少步采样器后仍然通过。 02. 最小可用理解 第一,把网络拆成每步仍要计算的部分,以及准备复用的昂贵部分。第二,保存上次完整计算得到的特征或残差,用便宜的变化指标决定是否刷新。第三,即使命中缓存,当前输入仍然要参与输出,采样器仍然继续更新状态。 以残差形式为例,设当前嵌入是 $h_i$,完整模块输出为 $G(h_i,t_i,c)$。缓存保存的是 $R_i=G(h_i,t_i,c)-h_i$。如果上次刷新发生在 $r$,当前近似模块输出是 $h_i+R_r$。符号 $i$ 表示第几次网络调用,$t_i$ 是这次调用的时间条件,$c$ 是固定的文本等条件;$r$ 是缓存的生成时刻,通常小于当前 $i$。 这三句话的重点是“当前输入加旧残差”。如果直接把整个旧输出拿回来,连当前输入的变化也被抹掉了。残差复用仍然有误差,只是保留了输入的直通更新。DeepCache 的 U-Net 高层特征缓存并不等于这个残差写法,不能把所有缓存方法都称为同一个算法。 03. 数学推导 先定义缓存误差发生在哪里 设 $E$ 是便宜的输入嵌入,$G$ 是昂贵模块,$H$ 是输出头。完整调用写为: $$h_i=E(x_i,t_i,c),\quad R_i=G(h_i,t_i,c)-h_i,\quad y_i=H(h_i+R_i,t_i,c)$$ $x_i$ 是采样器当前的潜变量;$h_i$ 是网络内部的嵌入,不必与 $x_i$ 形状相同;$R_i$ 是昂贵模块对嵌入的修正;$y_i$ 是送回采样器的网络预测。这里的预测可以按模型定义代表噪声或速度,缓存逻辑本身不替模型选择参数化。等式只是代数分解,不要求模型专门训练一个名叫“残差缓存”的层。 命中缓存后,输入嵌入和输出头仍按当前条件计算: $$\widehat{y}_i=H(h_i+R_r,t_i,c),\quad r<i$$ 帽子表示近似计算。比较精确与近似时,先固定当前同一个 $h_i$:误差来自 $R_r$ 替代了 $R_i$,而不是来自状态已经分叉。若输出头在所讨论的局部区域满足 Lipschitz 条件,常数记作 $L_H$,那么: $$\|\widehat{y}_i-y_i\|\leq L_H\|R_r-R_i\|$$ Lipschitz 条件的意思是输入变化不能被这个局部映射无限放大;$L_H$ 是放大上限。本文没有测出真实模型的 $L_H$,这个式子是带假设的分析工具,不是线上质量保证。它告诉我们,缓存内部误差小仍然需要考虑输出头的敏感性。 为什么使用便宜的变化代理 直接计算 $\|R_i-R_r\|$ 能判断缓存是否陈旧,但得到 $R_i$ 就已经做了昂贵计算。为了节省计算,需要先观察一个便宜的量 $m_i$。它可以由当前嵌入与时间条件构造。相邻调用的相对 L1 变化写为: $$d_i=\frac{\operatorname{mean}|m_i-m_{i-1}|}{\max(\operatorname{mean}|m_{i-1}|,\varepsilon)}$$ $\operatorname{mean}$ 对张量所有元素取平均,绝对值逐元素计算,$\varepsilon$ 是防止除零的小正数。两次张量形状一致时,均值之比等于 L1 范数之比。这个值没有单位,反映变化相对之前幅度有多大。分母很小时,它也可能变得敏感,不能把这种数值问题误判成内容剧烈变化。 本文教学实现取 $\varepsilon=10^{-12}$。该保护和后面非负累计属于教学实现的防御性处理;不能据此声称官方代码逐字采用相同逻辑。 最直觉的想法是:如果每一步变化都很小,就一直复用。但缓存对应的是刷新时刻 $r$,不是上一时刻。利用三角不等式: $$\|m_i-m_r\|\leq\sum_{j=r+1}^{i}\|m_j-m_{j-1}\|$$ 这里 $j$ 遍历刷新之后的调用。即使单步变化很小,连续多步也会积累成明显差异,因此策略应保留“自上次刷新以来的变化预算”。这个式子约束绝对变化;把每项分别除以不同的幅度后,不能直接称相对变化之和为严格误差界。 实际决策可以采用累计代理: $$A_i=A_{i-1}+g(d_i),\quad \text{refresh if } A_i\geq\delta$$ $A_i$ 是累计预算,$g$ 将输入变化映射到预估的输出变化,$\delta$ 是刷新阈值。完整计算后把预算归零并保存新残差。TeaCache 使用多项式重标定来改善这种代理;拟合关系是经验估计,不会把统计相关性变成对所有样本成立的数学上界。TeaCache 原文第 3.2–3.3 节 给出了指标与累计决策。 本文取 $g(d)=d$,以便只观察状态机。阈值 0.16 因而只属于本实验的代理尺度,不能迁移成真实模型的推荐阈值。换掉嵌入幅度、时间调制或采样路径,同一个阈值对应的刷新频率就会改变。 一次局部近似如何影响后续轨迹 现在让两条轨迹分别使用精确和近似输出。为便于推导,考虑显式 Euler 更新: $$x_{i+1}=x_i+\Delta s_i v(x_i,s_i),\quad \widehat{x }_{i+1}=\widehat{x }_i+\Delta s_i\widehat{v}(\widehat{x }_i,s_i)$$ $s_i$ 是积分坐标,$\Delta s_i>0$ 是步长,$v$ 是精确向量场,$\widehat{v}$ 使用缓存近似。令 $e_i=\|\widehat{x }_i-x_i\|$ 表示轨迹差异,令 $\eta_i=\|\widehat{v}(\widehat{x }_i,s_i)-v(\widehat{x }_i,s_i)\|$ 表示在近似轨迹当前状态上的局部缓存误差。 先在更新差值中加减 $v(\widehat{x }_i,s_i)$,把“缓存误差”和“输入状态不同”分开。再使用三角不等式,并假设 $v$ 对状态的 Lipschitz 常数为 $L$,得到: $$e_{i+1}\leq(1+\Delta s_iL)e_i+\Delta s_i\eta_i$$ 若初始状态相同,则 $e_0=0$。重复代入可得: $$e_N\leq\sum_{i=0}^{N-1}\Delta s_i\eta_i\prod_{j=i+1}^{N-1}(1+\Delta s_jL)$$ 空乘积取 1。这个展开说明早期误差会通过后续更新传播;最后一次强制刷新只能去掉最后一次局部缓存误差,不能恢复此前已经走偏的轨迹。真实系统可能局部收缩,实际误差小于这个上界;高阶求解器还涉及历史预测,不能把 Euler 的式子直接冒充所有采样器的误差定理。 为什么命中率不能直接当加速比 令 $N$ 是总调用数,$K$ 是完整计算数,$C_f$ 是完整调用成本,$C_h$ 是缓存命中调用成本,$C_o$ 是其他固定开销。则一个简单的成本模型是: $$S=\frac{NC_f+C_o}{KC_f+(N-K)C_h+C_o}$$ 命中调用仍要嵌入、判定、读缓存、运行输出头和更新采样器,$C_h$ 不会自动等于零。缓存张量跨设备读取甚至可能很贵。端到端还包括文本编码、VAE 解码和输出处理,跳过 Transformer 并不会同时省掉这些部分。 本文把 $C_f=1$、$C_h=0.08$、$C_o=0$ 当成教学假设。按实跑统计的 $K$ 代入可以算成本比,但没有用这些假设伪装 GPU 计时。部署时需要测出真实成本,再决定缓存是否值得加入。 04. 代码实现 完整脚本是文末的 cache_demo.py,依赖 NumPy 与 matplotlib,无需模型权重。它构造一个四层非线性残差函数,用固定随机种子产生输入和矩阵。输入形状是 $(B,L,D)=(1,4,16)$:一个批次、四个位置、十六个通道。张量使用 float64,目的是让读者稳定复核误差,不模拟混合精度推理。 脚本里的 cheap_input 对应 $E$,residual 对应 $R$,proxy 对应 $m$,budget 对应 $A$,threshold 对应 $\delta$。输出头取恒等映射。每次用当前 $h$ 加缓存残差得到预测,然后执行状态更新。时间标签从 1 递减到 0,而状态更新的步长固定为 $1/N$;这只是合成系统的定义,没有冒充 DDIM 或真实扩散调度器。 决策的核心代码如下;变量准备、精确对照、绘图及断言都在完整附录中: budget += max(change, 0.0) boundary = i in (0, STEPS - 1) calculate = boundary or budget >= threshold if calculate: saved_residual = residual(h, t) budget = 0.0 output = h + saved_residual 首次调用没有缓存,所以必须完整算。最后一次强制刷新属于本实现的策略。previous_proxy 每一步都更新,saved_residual 只有刷新时更新:把这两件事混淆,就会把“相邻变化累计”写成另一种算法。simulate 每次创建新的状态,避免后一条请求接着用前一条请求的残差。 为了画误差曲线,脚本还在每个步骤额外执行一次精确函数作为离线诊断。这些额外执行没有算进策略的 full_calls;因此实际运行这个带诊断的脚本并不是加速测速。计数回答的是策略本来会执行多少次昂贵计算,图回答的是当前近似轨迹上的局部输出误差有多大。 真实运行环境为 Python 3.11.9、NumPy 2.4.3。基线终点的元素范围为 [-0.791508, 0.614856]。输出如下: shape=(1, 4, 16), steps=40, seed=7 baseline_range=[-0.791508, 0.614856] name full_calls hits endpoint_relative_error assumed_speed_ratio exact 40 0 0.000000 1.000 adaptive_0.08 18 22 0.004939 2.024 adaptive_0.16 11 29 0.008222 3.003 adaptive_0.32 7 33 0.026326 4.149 uniform_4 11 29 0.008672 3.003 All assertions passed; figure=figures/cache_schedule.png 终点相对误差的定义是近似终点与精确终点之差的 L2 范数,除以精确终点的 L2 范数。因此 0.008222 是这个定义下约 0.8222% 的差异,不是 VBench 下降,也不是图像中有 0.8222% 的像素出错。最后一列来自前节假设的成本模型,没有真实延迟单位。 图上半部分每个方块表示一次完整残差计算,空缺表示复用;下半部分是离线诊断得到的局部输出 L2 误差。请先观察刷新是否集中在同一段,再看复用之间误差如何变化。曲线没有用感知质量指标,纵轴不能解读成画面可接受程度。配图由附录脚本生成,数据与上述输出来自同一次实验。 阈值 0.16 与固定间隔 4 都执行了 11 次完整计算,但终点误差略有差异。这是“同样预算也会因刷新位置不同而得到不同结果”的一个实例。它只有一个种子和一个人工系统,不能证明自适应策略在所有输入上优于固定间隔,更不能作为统计显著性结论。 05. 工业级实现对照 DeepCache 原文第 3.3 节 利用 U-Net 的分支结构:当前浅层特征继续更新,较深分支的高层特征从缓存取回并拼接。它缓存的是特定网络位置的特征,并非本文抽象的整段残差。官方 horseee/DeepCache 的 deepcache.py 提供 forward 包装与状态复位;对照时应先找缓存边界,再判断到底省掉了哪些模块。 TeaCache 的 HunyuanVideo 示例 中,teacache_forward 用时间调制后的输入计算变化代理,累计重标定后的变化,并保存 Transformer 主体的输入输出残差。首尾调用完整计算,命中后把旧残差加到当前图像嵌入上;输出层仍然运行。这里对应的是 2026-10-05 实际取回的源码,后续上游重构需要重新核对。 教学版与上述示例的差异很明确:教学版没有文本分支、真实 patchify、注意力或权重,也没有拟合官方多项式;它只验证预算与残差状态机。真实代码还可能把张量变化转成 CPU 标量,这种同步的代价不能从“指标计算量很小”中自动排除。需要把判定开销也纳入测量。 工程上,缓存状态应属于一次请求,而不是一个可被并发请求共享的随意全局变量。prompt、负向条件、guidance、分辨率、帧数、模型权重或调度器变化时,复用旧状态没有语义依据。只检查张量形状远远不够;两个请求可以形状相同、内容完全不同。失败中断后重新开始生成,也要建立新的缓存生命周期。 若模型采用双分支 CFG,必须考虑条件与无条件分支的缓存分别有效。按 $y=y_u+w(y_c-y_u)$ 写,$y_u$ 与 $y_c$ 分别是无条件和条件预测,$w$ 是 guidance 系数。预测误差满足: $$\|\Delta y\|\leq |w|\|\Delta y_c\|+|1-w|\|\Delta y_u\|$$ 因此较大的 guidance 可能放大分支近似误差。这个式子是代数上的工程分析,不声称所有视频模型都执行双分支 CFG;有些采用蒸馏后的 guidance 条件。接入前要先确认实际预测路径,再决定状态如何隔离。 还要区分采样步和网络调用。调度器可能在一个名义步里求值多次,或重访相同的时间标签。缓存的计数器应跟随真实调用语义。如果只拿界面上显示的步数当数组长度,很容易在错误的位置刷新,或者提前把状态清零。应记录每次实际调用的时间标签、是否刷新和缓存来源,再与调度器对齐。 06. 代价与边界 省的是昂贵模块重复计算,增加的是缓存内存、生命周期管理、判定开销和近似误差。保存一个形状为 $(B,L,D)$ 的残差张量,最低存储量为 $BLDb$ 字节,其中 $b$ 是每个元素的字节数;克隆的输入、多个分支及其他缓存还会额外占空间。长视频中 token 数很大,省算力并不自动省显存。 阈值没有通用刻度。代理的归一化方式、拟合数据、模型层、时间条件与采样器都影响其意义。在某模型上有效的多项式可能在另一个模型或调度器上失配。应在目标分布上先记录输入代理与真实残差变化的对应关系,再评估误判,不能只沿用别人仓库里的常数。 端点刷新是一种保守策略,不是一种证明。即使最后一次完整计算,网络读到的潜变量也已经经过缓存轨迹;它不能把积累误差自动洗掉。判断质量时要看整段生成结果,尤其是动作连续性、物体身份、细节与提示词对应关系,而不是只看最后一个调用的缓存误差归零。 已有很少采样步的蒸馏模型可能没有多少安全的复用空间。短轨迹上每次预测都更重要,进一步减少完整网络调用可能迅速损害质量。动态改变条件的交互式生成、频繁切换控制信号的任务,也不应默认共享同一段缓存。先建立可靠基线,才能判断具体组合是否值得部署。 评测时固定权重、prompt、种子、分辨率、帧数、采样器和 guidance,分别测完整网络计算时间与端到端延迟。预热、设备同步、重复测量和尾延迟都要说明;不要把首次编译的基线与已预热的缓存版放在一起比较。质量样例应包含快速运动、镜头切换、文字、遮挡和多主体交互等容易暴露近似问题的内容。 本文给出的误差曲线只能帮助理解缓存的数值行为。没有真实模型权重,没有样本集,也没有测量硬件延迟。因此本文不提供生产阈值,不给出“无损”的结论,也不声称这些 toy 误差能预测真实视频的主观质量。 还要分别检查代理误判的两种代价。代理高估变化会增加完整计算,主要损失效率;代理低估变化会继续使用陈旧残差,主要增加质量风险。这两种错误不能只用一个平均相关系数概括。即使整体相关性很高,少数快速变化的步骤被漏掉,也可能成为最明显的失败样例。应保留逐步记录,把异常视频定位到具体刷新区间,再判断是代理不敏感、阈值过大还是缓存位置选错。 为避免测试集上的阈值被过度优化,可以先用一组提示词选择阈值,再在另一组提示词上固定参数评估。报告平均质量以外,还要记录最差样例以及变化最大的场景。缓存有时能带来某项评分的小幅上升,但近似轨迹改变本身并不保证改进;应通过更多独立样本判断这是随机波动还是稳定效果。不能挑出一段更好看的视频,就把整条路线称作质量提升技术。 缓存决策在不同请求中可能产生不同的完整计算次数,所以服务端延迟也会随内容变化。只报告平均命中率,会遮蔽慢请求的行为;应同时关注完整计算次数的分布和尾延迟。如果业务要求固定延迟,可以限制最长复用间隔或最低刷新次数,但这些额外约束需要重新纳入成本账本与质量验证。 07. 经典论文脉络 DDIM:Denoising Diffusion Implicit Models,2010.02502。它提供理解不同采样路径的基础;跨步缓存要先说明是否保留原有离散调用,再比较误差,不能把省模块与改步长混为一谈。 DeepCache,2312.00858。它把 U-Net 高层特征的跨步复用转成免训练加速方案,主要启发是选择合适的结构边界,而不是对所有中间张量一概缓存。 Pyramid Attention Broadcast,2408.12588。它把视频生成中的注意力复用纳入设计,提示我们不同模块可以有不同刷新节奏。本文只把它作为技术脉络,不复现其性能结果。 TeaCache,2411.19108。它用时间调制输入的变化代理和重标定来选择刷新时刻,把问题从固定周期推进到输入感知的缓存决策。 论文数值也要带条件。DeepCache 的摘要报告 SD v1.5 上 2.3 倍加速与 CLIP Score 下降 0.05;这是作者特定设置下的结果,不能泛化成所有扩散模型都同样受益。原始摘要 支持这组数字。 TeaCache v2 的表 1 中,OpenSora-Plan、65 帧、512×512 的基线延迟为 99.65 秒,slow 配置为 22.62 秒,对应 4.41 倍;VBench 从 80.39% 到 80.32%,差值是 0.07 个百分点。fast 配置在同表报告更高加速,但质量也有不同取舍,所以本文引用 slow 配置,避免把不同配置的最优速度与最优质量拼成一个结果。原文表 1 可直接复核。 这些工作属于相关路线,没有一个统一的阈值或通用缓存张量。先检查缓存边界,再看刷新规则,最后检查采样路径是否改变,是比较它们时更稳定的顺序。方法名称相似、都声称免训练,不代表工程实现可以直接互换。 08. 常见误解 “跨步缓存与 KV cache 是同一种精确复用。” 自回归 KV cache 通常复用固定前缀的计算结果,而扩散过程中潜变量和时间条件在变化。这里复用的是可能陈旧的近似特征,必须单独评估误差。是否等价还取决于模型路径,不能仅凭有个 cache 字段就判断。 “相邻特征相似,所以可以一直复用。” 相邻相似不等于相隔很多步仍相似,更不等于输出质量不变。累计预算关注缓存年龄,刷新机制关注何时恢复完整计算。二者缺一不可。 “阈值 0.16 表示允许 16% 的质量损失。” 阈值只是累计代理的比较值。它不是终点误差、FID、VBench 或主观质量的直接比例。本实验阈值与官方多项式的阈值甚至不在相同刻度上。 “缓存命中后跳过了采样器更新。” 本文所有版本都更新状态 40 次。省掉的是一部分网络计算,当前输入与输出头仍然可以变化。如果代码命中后直接跳过调度器,它执行的是另一种算法。 “最后一次完整计算能补回前面所有误差。” 最后一次局部预测变准确,不等于输入轨迹回到基线。误差传播推导和实验都把这两件事区分开。 “调用次数少三倍,就是端到端快三倍。” 命中步骤、其他固定模块、同步与内存操作仍有成本。本文最后一列只是显式假设的成本比,实际脚本还带额外离线诊断,因此不能拿它的运行时间当加速报告。 09. 动手验证 先把附录完整代码保存为 cache_demo.py,准备 Python 3.10 以上环境,执行: python -m pip install numpy matplotlib python cache_demo.py 脚本按照所在目录生成实验结果文件和 figures/cache_schedule.png。数值应与第 04 节一致;换 NumPy 或底层数值库时,浮点尾数可能略有不同。比对完整计算次数、刷新时刻和误差到所列精度,比要求截图或图片压缩字节完全一致更合适。 首先把阈值设成 0。每个步骤都会刷新,完整计算次数应等于 40,终点应与精确基线逐元素一致。附录已经用断言验证这个条件;如果失败,要先修缓存状态机,不要调阈值掩盖问题。 其次比较 0.08、0.16 和 0.32。在本次固定实验里,完整计算次数逐步减少,终点误差逐步增加。这是一个具体样本的趋势,不是所有系统的单调性定理:刷新位置改变、误差相互抵消、轨迹分叉,都可能让某个中间阈值的结果反而更接近基线。 再比较 adaptive_0.16 与 uniform_4。两者的成本模型使用相同的 11 次完整计算,但时刻并不相同。观察图中误差峰值,不只盯终点一个数。如果想研究预算分配,固定计算次数,再调整刷新位置,比仅仅换一个阈值更容易把原因看清。 最后修改种子或残差函数中时间变化的频率,重新运行整套对照,并保留新的输出。这一操作没有本文预先承诺的性能结果;它用来观察原先代理是否仍能判断残差变化。若输入代理很平缓、真实残差突然变化,就能构造“指标漏报”的反例。把反例加入测试,比只展示一条顺利的曲线更有价值。 实际接入模型之前,还应验证两次独立请求是否有相同基线行为、错误中断是否清理状态、同形状不同条件是否互相污染。本文实验用每次 simulate 独立初始化来验证复位;真实服务还需要针对并发与重入的检查。 10. 延伸阅读 先从 DDIM 与高阶采样器 理解调用路径,再用 DiT 找嵌入、主体和输出头的边界,最后通过 性能建模与 Profiling 实测哪些部分值得省。 可以对照 KV cache 理解缓存有效性,也可继续读 量化 比较另一种减少推理成本的方法。知识树中的少步蒸馏节点仍为 planned;它是下一条相关学习路线,尚无已发布链接。跨步缓存、量化与蒸馏组合时,应重新建立质量和延迟基线。 附录:完整代码 09 节用到的脚本全文如下(cache_demo.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 cache_demo.py """Deterministic residual-cache toy; Python 3.10+, pip install numpy matplotlib. Run: python cache_demo.py No trained diffusion weights; cost ratios below are accounting assumptions. """ from pathlib import Path import json import platform import numpy as np import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt STEPS = 40 SHAPE = (1, 4, 16) SEED = 7 rng = np.random.default_rng(SEED) X0 = rng.normal(size=SHAPE).astype(np.float64) W = rng.normal(size=(16, 16)) / 8.0 V = rng.normal(size=(16, 16)) / 8.0 def cheap_input(x, t): return x + 0.05 * np.sin(2 * np.pi * t) def proxy(h, t): return h * (1 + 0.3 * np.sin(2 * np.pi * t)) + 0.1 * np.cos(2 * np.pi * t) def residual(h, t): z = h for _ in range(4): z = np.tanh(z @ W) return 0.15 * h + 0.4 * (z @ V) + 0.08 * np.sin(6 * np.pi * t) def simulate(threshold=None, period=None): x = X0.copy() previous_proxy = None saved_residual = None budget = 0.0 refreshed, local_errors = [], [] for i in range(STEPS): t = 1.0 - i / (STEPS - 1) h = cheap_input(x, t) m = proxy(h, t) change = 0.0 if previous_proxy is None else float( np.mean(np.abs(m - previous_proxy)) / max(np.mean(np.abs(previous_proxy)), 1e-12)) budget += max(change, 0.0) boundary = i in (0, STEPS - 1) if threshold is None and period is None: calculate = True elif period is not None: calculate = boundary or i % period == 0 else: calculate = boundary or budget >= threshold if calculate: saved_residual = residual(h, t) budget = 0.0 output = h + saved_residual # Offline diagnostic only; extra evaluation is excluded from work accounting. exact = h + residual(h, t) local_errors.append(float(np.linalg.norm(output - exact))) refreshed.append(bool(calculate)) previous_proxy = m.copy() x = x - output / STEPS return {"x": x, "refreshed": refreshed, "local_errors": local_errors} def experiment(): baseline = simulate() rows = [] results = [] settings = [("exact", None, None), ("adaptive_0.08", 0.08, None), ("adaptive_0.16", 0.16, None), ("adaptive_0.32", 0.32, None), ("uniform_4", None, 4)] for name, threshold, period in settings: result = simulate(threshold, period) calls = sum(result["refreshed"]) relative_error = float(np.linalg.norm(result["x"] - baseline["x"]) / np.linalg.norm(baseline["x"])) # Full call = 1 unit; cache hit = 0.08 unit, solely for illustration. assumed_work = calls + (STEPS - calls) * 0.08 rows.append({"name": name, "full_calls": calls, "hits": STEPS - calls, "relative_endpoint_error": relative_error, "assumed_speed_ratio": STEPS / assumed_work, "refresh_indices": [i for i, flag in enumerate(result["refreshed"]) if flag]}) results.append(result) zero = simulate(threshold=0.0) assert np.array_equal(zero["x"], baseline["x"]) assert sum(zero["refreshed"]) == STEPS assert all(r["refreshed"][0] and r["refreshed"][-1] for r in results) assert np.array_equal(simulate(threshold=0.16)["x"], results[2]["x"]) assert all(np.isfinite(r["x"]).all() for r in results) return rows, results, baseline def make_figure(rows, results, destination): fig, axes = plt.subplots(2, 1, figsize=(11, 7), constrained_layout=True) colors = ["#2563eb", "#ea580c", "#16a34a", "#dc2626", "#9333ea"] for y, (row, result) in enumerate(zip(rows, results)): idx = np.flatnonzero(result["refreshed"]) axes[0].scatter(idx, np.full(len(idx), y), marker="s", s=38, color=colors[y]) axes[0].set_yticks(range(len(rows)), [r["name"] for r in rows]) axes[0].set_xlim(-1, STEPS) axes[0].set_xlabel("Sampling call index (0-based)") axes[0].set_title("Squares = full residual computation; gaps = cache reuse") for y, (row, result) in enumerate(zip(rows[1:], results[1:]), start=1): axes[1].plot(result["local_errors"], label=row["name"], linewidth=1.7, color=colors[y]) axes[1].set_xlabel("Sampling call index (0-based)") axes[1].set_ylabel("Local output L2 error") axes[1].set_title("Offline diagnostics on each cached trajectory; no quality metric") axes[1].legend(ncol=2, fontsize=9) axes[1].grid(alpha=0.25) fig.suptitle("Residual caching: refresh decisions and approximation error", fontsize=14) fig.savefig(destination, dpi=150) plt.close(fig) if __name__ == "__main__": rows, results, baseline = experiment() folder = Path(__file__).resolve().parent.parent (folder / "figures").mkdir(parents=True, exist_ok=True) make_figure(rows, results, folder / "figures" / "cache_schedule.png") record = {"seed": SEED, "steps": STEPS, "shape": list(SHAPE), "python": platform.python_version(), "numpy": np.__version__, "baseline_min": float(baseline["x"].min()), "baseline_max": float(baseline["x"].max()), "results": rows, "assertions": "zero-threshold exactness, endpoints, request reset, finite outputs passed"} (folder / "experiment_results.json").write_text( json.dumps(record, ensure_ascii=False, indent=2) + "\n", encoding="utf-8") print(f"shape={SHAPE}, steps={STEPS}, seed={SEED}") print(f"baseline_range=[{record['baseline_min']:.6f}, {record['baseline_max']:.6f}]") print("name full_calls hits endpoint_relative_error assumed_speed_ratio") for row in rows: print(f"{row['name']} {row['full_calls']} {row['hits']} " f"{row['relative_endpoint_error']:.6f} {row['assumed_speed_ratio']:.3f}") print("All assertions passed; figure=figures/cache_schedule.png")
2026年10月05日
4 阅读
0 评论
0 点赞
2026-10-04
AIGC 基本功知识树
每日论文速读解决的是广度——今天世界上发生了什么。但读论文有个前提:你得先看得懂。 这里是另一条线:把视觉生成的底层零件一个个拆开讲透。每篇都给数学推导 + 能跑的代码 + 经典论文出处,并标注前置知识,你可以顺着依赖链一路读下来。 当前规划 38 个知识点,已发布 33 篇。 图例:● 已发布 · ◍ 已发布待更新 · ◐ 正在写 · ○ 计划中 数学与优化基础(2/2) 概率、变分推断、随机微分方程——读懂扩散模型公式的最小前置集合。 ● 变分下界与重参数化(ELBO) · 难度:入门前置 不理解 ELBO 就没法理解 VAE 的 KL 项为什么要加权,也没法理解扩散模型的训练目标从哪来。 ● 扩散过程的前向与反向推导(SDE) · 难度:入门前置 · 前置:ELBO 把 DDPM 的离散公式和 SDE 的连续视角对上,后面所有采样器的差异都能一句话解释。 注意力与核心零件(5/7) 注意力、位置编码、稀疏激活与循环记忆——视频生成里最吃显存、最难外推也最影响长程建模的部件。 ● 自注意力机制的计算与显存账本(MHA) · 难度:核心必修 先把 O(N²) 的常数项算清楚,才能判断后面各种稀疏/线性方案到底省在哪一项。 ● 旋转位置编码 RoPE 的原理与实现(RoPE) · 难度:核心必修 · 前置:MHA 许多现代视频 DiT 使用 RoPE;理解旋转和相对位置性质,才能区分坐标可计算与生成质量可外推。 ● FlashAttention 为什么不需要存下注意力矩阵(FlashAttn) · 难度:进阶 · 前置:MHA online softmax 这一个技巧撑起了整个长序列时代,值得把递推式一步步推一遍。 ● 视频 DiT 里的 3D RoPE 与分辨率外推(3D-RoPE) · 难度:进阶 · 前置:RoPE 换分辨率或帧数时,3D RoPE 的坐标、频率与轴分配是需排查的因素之一,还要检查训练分布、VAE 与注意力实现。 ● 混合专家 MoE:稀疏激活怎么省算力(MoE) · 难度:工程实战 · 前置:MHA 统一多模态模型开始普遍用 MoE 扛参数量,但路由不均衡带来的训练不稳定很少被讲清楚。 ○ 线性、循环与记忆架构(LinearMem) · 难度:前沿 · 前置:MHA 长序列架构正在从保存完整注意力矩阵转向可更新状态和深度循环;把复杂度、记忆容量与并行性放进同一框架,才能判断它们何时真能替代 Softmax Attention。 ○ 视频生成里的稀疏注意力(SparseAttn) · 难度:前沿 · 前置:FlashAttn、3D-RoPE 视频长序列的全注意力成本很高;稀疏化可以减少计算,但免训练部署的质量取舍需和精确 IO 优化、缓存等方案比较。 生成范式(7/9) 从 DDPM 到流匹配,以及把 50 步压到 4 步的蒸馏路线。 ● DDPM 训练目标与采样流程(DDPM) · 难度:核心必修 · 前置:SDE DDPM 是理解扩散模型的重要起点;先读懂其 loss 与采样循环,再比较其他生成路径和训练目标。 ● 分类器无关引导 CFG 的代价与调法(CFG) · 难度:核心必修 · 前置:DDPM CFG 让每步算两遍,是推理成本里最容易被忽视的 2×,也是蒸馏首先要干掉的对象。 ● 从 DDIM 到高阶采样器(DDIM) · 难度:进阶 · 前置:DDPM 采样器换一个、缓存策略就得重调——这是推理加速最常见的踩坑点。 ● 流匹配与 Rectified Flow(FlowMatching) · 难度:进阶 · 前置:SDE、DDPM 流匹配被许多近期生成模型采用;它与 DDPM 的概率路径、监督目标及采样方式值得系统比较。 ● 潜空间扩散与 Stable Diffusion 架构(LDM) · 难度:进阶 · 前置:DDPM、VAE 把扩散搬进潜空间这一步,直接决定了今天视觉生成的算力可行性。 ● DiT:用 Transformer 替掉 UNet(DiT) · 难度:进阶 · 前置:LDM、MHA adaLN-Zero 这个小设计是 DiT 能稳定训起来的关键,值得逐行对照公式看。 ● 自回归视频生成与 Forcing 范式(Forcing) · 难度:工程实战 · 前置:DDPM、DiT 双向扩散没法流式出帧,Forcing 系列是把视频生成变成可交互的关键一步,也是当下最活跃的范式。 ○ 少步蒸馏:从 50 步到 4 步(Distill) · 难度:前沿 · 前置:DDIM、CFG 蒸馏是过去两年推理加速收益最大的一条线,也是最容易把质量搞崩的一条。 ○ 世界模型:从视频生成到可交互环境(WorldModel) · 难度:前沿 · 前置:Forcing 视频生成正在从「出片」转向「可交互环境」,这是范式转变而不是又一个 SOTA 分数。 表征与压缩(4/4) VAE / Tokenizer——视频生成的「地基」,决定了上限和伪影形态。 ● VAE 结构与训练目标(VAE) · 难度:核心必修 · 前置:ELBO 潜空间扩散依赖编码器输出的统计尺度;像素空间扩散不使用 VAE,不能将 scaling_factor 推广到所有扩散模型。 ● 离散化表征:VQ-VAE 与 VQGAN(VQGAN) · 难度:进阶 · 前置:VAE 离散自回归视觉模型依赖 tokenizer,码本利用率是重要问题;连续潜变量自回归路线不适用这一前提。 ● 视频 VAE 的时空压缩结构(VideoVAE) · 难度:工程实战 · 前置:VAE 时间维压缩比和因果卷积的实现方式,直接决定了长视频能不能逐块解码而不接缝。 ● 视频 VAE 的常见 loss 组合(VAELoss) · 难度:工程实战 · 前置:VideoVAE L1 + KL + LPIPS + GAN 四项的权重配比是玄学重灾区,把每项在管什么讲透很有价值。 对齐与强化学习(4/4) 从 PPO 到 GRPO,以及怎么把 RL 用到扩散和视频生成上。 ● 策略梯度与 PPO 基础(PPO) · 难度:核心必修 PPO 是在线策略优化的重要基础;掌握优势估计和裁剪后,再区分 GRPO 与直接偏好优化等不同路线。 ● 从 DPO 到 GRPO:去掉价值网络(GRPO) · 难度:进阶 · 前置:PPO GRPO 用组内相对奖励省去 critic;净成本还取决于组采样、参考模型、奖励评估和显存,不能统一声称降低一个量级。 ● 把 RL 用到扩散模型上(DiffusionRL) · 难度:前沿 · 前置:GRPO、DDIM 把多步去噪当成一条 MDP 轨迹,是理解视频 RL 各种做法的统一视角。 ● 视频生成中的强化学习与奖励模型(VideoRL) · 难度:前沿 · 前置:DiffusionRL、FlowMatching 视频的奖励要同时管画质、运动和指令遵循,reward hacking 在这里表现得最明显。 分布式训练(4/4) 参数、数据、序列三个维度怎么切,以及切完之后通信量变成多少。 ● 数据并行与 ZeRO 显存切分(ZeRO) · 难度:进阶 先把「参数 + 梯度 + 优化器状态」的显存账算清楚,才知道该切哪一部分。 ● 混合精度与数值稳定性(AMP) · 难度:进阶 BF16、FP16 与 FP8 的范围、精度和缩放机制不同,排查异常时要区分前向溢出、梯度下溢与舍入误差。 ● 张量并行与流水线并行(TP-PP) · 难度:工程实战 · 前置:ZeRO TP 的通信量和 PP 的气泡率都能手算,算完就知道并行度该怎么配。 ● 序列并行与 Ring Attention(SP) · 难度:前沿 · 前置:TP-PP、FlashAttn 超长视频序列可能需要序列或上下文并行;是否值得使用取决于单卡容量、通信和计算重叠。 推理加速与部署(5/5) 先用性能模型定位算力、带宽和显存瓶颈,再用缓存、稀疏、量化与算子融合把生成时间从分钟压到秒。 ● 性能建模与 Profiling:算力、带宽与显存账本(Roofline) · 难度:进阶 · 前置:AMP 不先判断算力、带宽还是显存容量在卡住系统,量化、融合、缓存很容易做成负优化;这篇提供所有推理优化节点共同的测量方法。 ● KV Cache 与自回归视频生成(KVCache) · 难度:进阶 · 前置:MHA、Roofline 自回归视频模型正在把 LLM 的这套缓存工程整体搬过来,值得先打好底。 ● 算子融合与 CUDA Graph(Fusion) · 难度:工程实战 · 前置:MHA、Roofline 小算子多的模型往往卡在访存和 launch 开销上,融合的收益比换算法更确定。 ● 量化:从 INT8 到 FP4(Quant) · 难度:工程实战 · 前置:AMP、Roofline 激活里的离群值是量化掉点的主因,SmoothQuant 的迁移技巧值得逐步推演一遍。 ● 扩散模型的跨步缓存复用(DiffCache) · 难度:前沿 · 前置:DDIM、DiT、Roofline 相邻去噪步的特征高度相似,这是免训练加速里性价比最高的一类手段。 评测与指标(2/3) FID/CLIP/VBench 各自测的是什么,以及它们什么时候会骗人。 ● FID / CLIP Score 到底测了什么(FID) · 难度:核心必修 FID 对样本量和预处理极其敏感,不同论文的数字经常根本不可比。 ○ 视频生成评测:VBench 与人工验收(VBench) · 难度:进阶 · 前置:FID 加速类工作最爱报 VBench 总分不变,但分维度看往往能发现明显退化。 ● 音频质量评测:MOS、PESQ 与 FAD 各测什么(AudioEval) · 难度:工程实战 音视频生成里音频质量几乎全靠几个数字说话,但这些数字测的根本不是一回事:SI-SDR 只管波形对齐、PESQ 只为通信语音设计、FAD 看的是分布而不是单条样本的保真度。选错指标会得到自欺欺人的结论——比如用波形指标验收 codec,数字很好听但高频毛刺全在。 接下来会写 线性、循环与记忆架构(LinearMem) —— 长序列架构正在从保存完整注意力矩阵转向可更新状态和深度循环;把复杂度、记忆容量与并行性放进同一框架,才能判断它们何时真能替代 Softmax Attention。 视频生成里的稀疏注意力(SparseAttn) —— 视频长序列的全注意力成本很高;稀疏化可以减少计算,但免训练部署的质量取舍需和精确 IO 优化、缓存等方案比较。 少步蒸馏:从 50 步到 4 步(Distill) —— 蒸馏是过去两年推理加速收益最大的一条线,也是最容易把质量搞崩的一条。 最后更新:2026-10-05 12:40
2026年10月04日
11 阅读
0 评论
3 点赞
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 点赞
2026-09-17
AIGC 基本功|把 RL 用到扩散模型上-DiffusionRL
把 RL 用到扩散模型上:从一张图的奖励到整条去噪轨迹的梯度 所属方向:对齐与强化学习 | 难度:前沿 | 前置知识:GRPO、DDIM 采样 关键词:DiffusionRL、DDPO、去噪轨迹、终点奖励、信用分配、DPOK、DRaFT 01. 为什么需要它 扩散模型通常先学习“像训练数据”,但产品真正关心的目标往往不是训练集似然:文字有没有正确渲染、人物是不是多了一根手指、构图是否讨喜、生成物能不能被某个下游系统识别、视频运动是否自然。这些目标可能来自人工打分、视觉语言模型、OCR、物理模拟器,甚至一段只能返回分数的业务程序。它们大多不可微,也不一定有成对偏好数据,却都能回答同一个问题:这张最终样本值多少分? 直接把最终图片当成一次普通动作,会丢掉扩散模型最重要的结构:图片不是一次生成的,而是从纯噪声开始,经过几十次随机去噪才得到的。假设采样有 50 步,最后的文字错了,到底是哪一步需要改?如果只对最终结果拟合,模型既看不到每一步在旧策略下出现的概率,也无法使用 PPO 的重要性比率约束更新;如果硬把 50 个概率相乘成一个轨迹比率,数值又很容易趋近 0 或爆炸。 DiffusionRL 的关键不是发明一种新的扩散网络,而是换一个观察角度: 把一次完整采样视为一条有限时域 MDP 轨迹;当前 latent 是状态,下一 latent 的随机采样是动作,最终图片的评分是终点奖励。 这样,奖励虽然只在最后出现,整条轨迹的 log-prob 仍然可以参与策略梯度。DDPO(Denoising Diffusion Policy Optimization)正是把这个视角变成了可训练算法:它允许我们用 JPEG 可压缩性、目标检测器、审美模型或人工反馈等黑盒奖励微调扩散模型,而不要求对奖励函数求导。 这也是理解视频生成 RL 的统一入口。图片去噪已经有“时间步 × latent”的信用分配;到了视频,还会再叠加帧间运动、时序一致性和长时奖励。先把图像版的轨迹与概率比率看清,后面的算法名字就不再是一堆孤立缩写。 02. 最小可用理解 给定提示词 $c$,采样从 $x_T\sim\mathcal N(0,I)$ 开始;每个去噪步骤都由模型定义一个策略 $p_\theta(x_{t-1}\mid x_t,c)$。 完整轨迹的概率是所有步骤概率的乘积,所以它的 log-prob 是逐步 log-prob 之和;终点奖励 $R(x_0,c)$ 因而能乘到每一步的 score function 上。 实际训练不直接做无约束 REINFORCE,而是像 PPO 一样保存旧策略的逐步 log-prob、重新计算新策略 log-prob、裁剪概率比率,并用基线或组内标准化降低方差。 图里有一个必须说准的细节:网络通常输出噪声、velocity 或转移均值的参数,但在 MDP 记号中,“动作”是实际采样出来的 $x_{t-1}$。也就是说,网络参数化动作分布,采样器从这个分布中抽出动作;不要把“网络输出”和“环境中发生的动作”混成一件事。 03. 数学推导 3.1 先把一次随机去噪写成概率分布 为避免被不同预测参数化淹没,先只保留采样器真正需要的形式。对当前 latent $x_t$,模型根据时间步 $t$ 和条件 $c$ 预测噪声,再由调度器算出一个转移均值: $$\mu_\theta(x_t,t,c).$$ 带随机性的 DDIM 步可以统一写成: $$x_{t-1}=\mu_\theta(x_t,t,c)+\sigma_t z_t,\qquad z_t\sim\mathcal N(0,I).$$ 因此条件转移是一个高斯策略: $$\pi_\theta(a_t\mid s_t)=p_\theta(x_{t-1}\mid x_t,t,c)=\mathcal N\!\left(x_{t-1};\mu_\theta(x_t,t,c),\sigma_t^2 I\right).$$ DDIM 中的随机强度由 $\eta$ 控制。训练 DDPO 时通常必须让 $\eta>0$,这样被纳入似然目标的随机步骤才有普通高斯密度可用于 log-prob 和重要性比率。注意 $\eta>0$ 仍不保证所有步骤 $\sigma_t>0$:落到干净端的终步可能方差为零,需排除该步或显式采用非零方差,不能直接算高斯 log-prob。若 $\eta=0$,转移退化为确定性映射,$\sigma_t=0$,直接套高斯 log-prob 会除以零;确定性采样仍可用于可微奖励反传,但不能不加处理地套用这套随机策略梯度。 3.2 MDP 的四个元素 把扩散采样改写成 MDP 时,可使用下面的对应关系: RL 概念 扩散采样中的对象 状态 $s_t$ 当前 latent、时间步和条件:$(x_t,t,c)$ 动作 $a_t$ 从模型给出的分布中采样下一 latent:$x_{t-1}$ 策略 $\pi_\theta$ 反向去噪转移 $p_\theta(x_{t-1}\mid x_t,t,c)$ 奖励 中间步骤通常为 0;解码 $x_0$ 后得到 $R(x_0,c)$ 这里的“环境”几乎没有额外动力学:动作 $x_{t-1}$ 本身就成为下一状态的一部分,时间步从 $t$ 减到 $t-1$。文本编码器、VAE 解码器和奖励模型通常不属于策略参数;真正更新的是 UNet、DiT 或其 LoRA 适配器。 3.3 终点一个分数,为什么每一步都有梯度 给定提示词,一条轨迹记作 $\tau=(x_T,x_{T-1},\ldots,x_0)$。因为初始噪声不依赖模型参数,轨迹概率可分解为: $$p_\theta(\tau\mid c)=p(x_T)\prod_{t=T}^{1}p_\theta(x_{t-1}\mid x_t,t,c).$$ 取对数后,乘积变为求和: $$\log p_\theta(\tau\mid c)=\log p(x_T)+\sum_{t=T}^{1}\log p_\theta(x_{t-1}\mid x_t,t,c).$$ 目标是最大化期望终点奖励: $$J(\theta)=\mathbb E_{c,\tau\sim p_\theta}[R(x_0,c)].$$ 应用 likelihood-ratio trick: $$\nabla_\theta J(\theta)=\mathbb E\!\left[R(x_0,c)\nabla_\theta\log p_\theta(\tau\mid c)\right]=\mathbb E\!\left[R(x_0,c)\sum_{t=T}^{1}\nabla_\theta\log p_\theta(x_{t-1}\mid x_t,t,c)\right].$$ 这就是“奖励只在终点,梯度覆盖全轨迹”的来源。实践中会减去不依赖当前动作的基线 $b(c)$,得到优势: $$A=R(x_0,c)-b(c).$$ 减基线不会改变期望梯度,却能显著降低采样方差。若同一个 prompt 一次生成多张图,可以用组内均值和标准差做归一化: $$A_i=\frac{R_i-\operatorname{mean}(R_{1:G})}{\operatorname{std}(R_{1:G})+\epsilon}.$$ 组均值包含自身时,未标准化的中心化 score 估计在独立采样条件下有 $(1-1/G)$ 的缩放;随机标准差还会改变各组权重。因此上面的无偏基线证明不能直接套到组内标准化。它和 GRPO 的“组内相对分数”很相似,不过动作空间从 token 变成了高维 latent 转移。 3.4 每一步的 log-prob 到底怎么算 设动作 $x_{t-1}$ 有 $D$ 个独立高斯坐标、固定方差为 $\sigma_t^2>0$。真实联合 log-prob 必须对坐标求和: $$\ell_\theta=-\frac{\|x_{t-1}-\mu_\theta\|^2}{2\sigma_t^2}-D\log\sigma_t-\frac D2\log(2\pi).$$ 真实逐步联合概率比为 $r_t=\exp(\ell_\theta-\ell_{old})$。DDPO 作者实现 则把逐坐标 log-prob 取平均,即 $\bar\ell=\ell/D$,再计算 $$\bar r_t=\exp(\bar\ell_\theta-\bar\ell_{old})=r_t^{1/D}.$$ 这是一种控制高维数值尺度的工程替代比率,不再是严格的联合重要性权重;它也改变了 clip 的尺度。本文一维玩具 $D=1$ 时二者相同,多维复现时必须明确采用哪个约定。下式的 $r_t$ 应按所选约定解释。 PPO 风格的裁剪目标是: $$L_t(\theta)=\min\!\left(r_t A,\operatorname{clip}(r_t,1-\epsilon,1+\epsilon)A\right).$$ 训练代码通常最小化 $-L_t$。关键是先按转移计算比率再裁剪,不是先把所有 $r_t$ 相乘成一个轨迹比率。后者会随步数快速积累方差,也让某一步的异常比率拖垮整条样本。 3.5 同一个优势,不代表每一步得到同样的梯度 优势 $A$ 可以广播给每一步,但 score function 不同。对固定方差的高斯策略,有: $$\nabla_\theta\log p_\theta(x_{t-1}\mid x_t,c)=\frac{x_{t-1}-\mu_\theta}{\sigma_t^2}\nabla_\theta\mu_\theta.$$ 时间步的噪声尺度、采样残差、网络 Jacobian 都不同,因此各步的梯度贡献不会相等。终点附近通常对最终结果影响更直接,但并不存在“最后一步天然包办所有学习”的定理;调度器与模型参数化会改变信用在各步之间的分布。 04. 代码实现 下面的玩具过程只有一个可训练标量 $\theta$,但保留了 DDPO 的核心结构:8 步高斯去噪、只在终点计算奖励、保存逐步 log-prob、用中心化优势做 REINFORCE,并在最后演示新旧策略的 PPO 比率。 状态转移是: $$x_{t-1}=0.72x_t+0.18\theta+0.45z_t,\qquad z_t\sim\mathcal N(0,1).$$ 终点目标为 1.5,奖励为: $$R(x_0)=-(x_0-1.5)^2.$$ 完整代码见文末。直接运行: python3 ddpo_minimal.py 真实输出如下: trajectory_shape= (4096, 9) transition_log_prob_shape= (4096, 8) before_terminal_mean= 0.0052 before_reward_mean= -2.6513 after_theta= 2.2677 after_terminal_mean= 1.3577 after_reward_mean= -0.4370 per_step_gradient= [0.0088 0.0085 0.014 0.024 0.0358 0.0435 0.0578 0.0883] gradient_sample_std_raw_centered= 1.6981 1.3899 ppo_ratio_mean_min_max= 1.0010 0.4051 2.2337 ppo_clip_fraction= 0.3108 history_first= [ 0. -0.0165 -2.7295] history_last= [ 2.2677 1.3601 -0.4503] 几个 shape 很重要:4096 条轨迹、每条 8 个转移,所以状态数组有 9 个点,而 log-prob 只有 8 个。训练前 $x_0$ 均值约为 0;40 次更新后移到 1.36,接近目标 1.5,平均奖励从 -2.65 提升到约 -0.44。奖励是负平方误差,所以“更大”意味着更接近 0,而不是数值绝对值更大。 中心化前后,单样本梯度估计的标准差从 1.70 降到 1.39。这只是最简单的 batch baseline;工业实现还会按 prompt 维护统计量,避免一个本来就容易拿高分的 prompt 主导更新。 最后把参数轻微扰动,逐步比率均值仍约为 1,但极端值已经到 0.41 和 2.23,有约 31% 的转移落到裁剪区间之外。这说明“平均比率看起来正常”不足以判断更新是否安全,还要看 clip fraction、KL 和比率分位数。 右图里越靠近终点的转移贡献越大,是这个线性玩具过程的收缩系数 0.72 造成的:早期扰动经过多次收缩后影响变小。真实扩散模型不会总呈现这种单调形状,但它直观说明了:即使每一步共享同一个终点优势,信用仍由每一步的局部概率结构决定。 05. 工业级实现对照 截至 2026 年 9 月,DDPO 论文作者维护的参考实现 kvablack/ddpo-pytorch 仍是理解训练循环最直接的代码入口。其 scripts/train.py 和带 log-prob 的 ddim_step_with_logprob 展示了完整数据流: 采样 rollout:从 prompt 和初始噪声出发,保存形状近似为 [batch, steps + 1, C, H, W] 的 latents,以及 [batch, steps] 的旧策略 log-prob。 只给终点打分:将 $x_0$ 经 VAE 解码成图片,用 JPEG 可压缩性、审美分数、目标检测、视觉语言模型或业务规则计算奖励。慢奖励可以并发执行,但返回时必须和样本一一对齐。 构造优势:对全批次奖励标准化,或用 per-prompt statistics tracker 做条件化标准化。同一条轨迹的优势广播到各去噪时间步。 打乱样本与时间步:实现不仅 shuffle 样本,还可对每个样本打乱时间步顺序,降低相邻去噪步相关性对小批次更新的影响。 重算 log-prob:把保存的 $(x_t,x_{t-1},t,c)$ 喂给当前模型,通过 patched DDIM step 得到新策略 log-prob,而不是重新采样一个 $x_{t-1}$。 PPO 裁剪更新:计算逐步新旧比率,应用 clipped surrogate loss;通常只训练 LoRA,以降低显存、通信和策略漂移。 这个实现细节解释了为什么训练内存明显大于普通推理:rollout 阶段要保存整条 latent 轨迹和旧 log-prob,更新阶段又要为当前时间步建立反向图。若 rollout batch 为 $B$、去噪步数为 $T$,一轮采样至少处理 $B\times T$ 次 UNet/DiT 前向;若每批 rollout 做 $K$ 个 inner epochs,更新成本还会近似再乘 $K$。梯度累积只改变峰值显存,不会凭空消除总计算量。 早期 Hugging Face TRL 版本曾提供 DDPOTrainer,但当前上游目录与公开 API 已发生重构。读旧教程时应锁定对应版本,不要假设 trl/trainer/ddpo_trainer.py 在最新版仍是稳定入口;本文因此把论文作者仓库作为代码锚点。 从 rollout 到 update 的数据契约 工程里最值得先写清的不是 Trainer 类,而是一条样本在两个阶段之间必须携带什么。rollout 结束后,每条样本至少要保存 prompt 或其编码、全部时间步、从初始噪声到终点的 latents、旧策略逐步 log-prob、最终图片、原始奖励和经过标准化的优势。若使用 classifier-free guidance,还必须固定无条件分支、guidance scale 和文本编码方式;否则更新阶段重算出来的均值并不是采样时那条策略的均值。 update 阶段只取保存的当前 latent 和下一 latent,调用当前模型重算同一个已发生动作的 log-prob。这里绝对不能重新抽噪声:PPO 问的是“新策略对旧动作给出多大概率”,不是“新策略这次随机采到了什么”。这个区别在 token PPO 中很直观——不会在更新时重采一句回答——换成连续 latent 后却很容易藏在调度器接口里。 还应给 rollout 批次分配不可变的 sample id,并把 prompt、随机种子、reward 版本、模型 checkpoint 和 scheduler 配置写进元数据。这样一旦某个 batch 的 ratio 或奖励异常,能够重放出同一条轨迹并定位问题。只记录最后的 PNG 不够,因为不同 latent 轨迹可能解码成肉眼相近的图片,而策略梯度依赖的是当时每一步的概率。 多卡训练时,全局优势标准化也必须基于所有进程聚合后的奖励,而不是每张卡各算各的。否则不同卡的 prompt 难度稍有偏差,就会得到不同的零点和尺度;随后再做梯度平均,相当于混合了多套不一致的目标。相反,若采用 per-prompt 统计器,应明确冷启动策略:某个 prompt 样本太少时回退到全局统计,避免标准差接近零导致优势被放大。 DPOK:为什么还要加 KL 纯奖励最大化很容易把模型推离预训练分布。DPOK(Diffusion Policy Optimization with KL regularization)在在线策略梯度中加入参考模型约束,可抽象为: $$J_{\text{DPOK}}(\theta)=\mathbb E[R(x_0,c)]-\beta\,D_{\mathrm{KL}}(p_\theta\|p_{\text{ref}}).$$ $\beta$ 越大,模型越保守;越小,越容易快速追逐奖励,也越容易丢失多样性或出现奖励投机。PPO clipping 约束的是一次更新相对旧策略的变化,reference KL 约束的是长期相对基座模型的漂移,两者作用并不相同。 DRaFT:奖励可微时,可以不走 score function 如果奖励模型对像素可微,且采样器的整条计算图可以保留,就能把梯度从奖励直接穿过 VAE 与采样步骤反传到扩散模型。DRaFT 将这一路线系统化,并提出截断反传等变体来控制内存与梯度不稳定。 两条路线的选择很清楚: DDPO / DPOK:奖励可以是完全黑盒,只需要返回标量;代价是策略梯度方差较高,需要大量 on-policy 样本。 DRaFT 类方法:能利用奖励的路径导数,梯度通常更直接;但奖励必须可微,且跨很多采样步骤保存或重算计算图,显存与稳定性是主要问题。 它们可以共享同一个奖励定义,却不能把“对奖励求导”和“用奖励乘 log-prob”混写成一套推导。 06. 代价与边界 6.1 它省下了什么 不需要奖励函数可微,外部 API、人工评分和离散程序都能接入。 不需要为每个目标重新制作大规模成对偏好数据;在线采样能探索当前策略真正会生成的区域。 能复用 PPO 的裁剪、KL、优势归一化和诊断工具,把更新幅度控制在可观察范围内。 6.2 它付出了什么 on-policy 成本高:策略一更新,旧 rollout 很快过期;数据复用能力弱于离线偏好优化。 终点信用粗糙:所有步骤共享终点优势,方差高;步数越多、latent 越高维,问题越明显。 奖励调用可能成为瓶颈:大型 VLM 或人工评分比去噪本身还慢,异步队列与缓存很重要。 探索与画质有冲突:增大 DDIM 的 $\eta$ 有利于随机策略探索,却可能改变原有采样质量与分布。 奖励投机不会自动消失:OCR 奖励可能鼓励巨大文字,审美分数可能压低多样性,检测器奖励可能诱导模型画出分类器偏爱的纹理。 6.3 上线前至少记录这些指标 不要只看 mean reward。最低限度还应记录: 奖励均值、标准差和分位数,按 prompt 类别拆分; 新旧策略的 approximate KL、ratio 分位数与 clip fraction; 每个时间步的 loss、KL 或梯度范数,检查是否只剩少数步在学习; 多样性指标和基础画质指标,防止奖励提高但模式坍缩; 人工盲评或独立奖励模型的 holdout 分数,防止只攻破训练用 reward; 无效图片、NaN、全黑/过曝和安全过滤命中率。 如果奖励模型本身正在更新,还要给 reward 版本打快照。否则同一条训练曲线纵轴含义会变化,前后 reward 数字不可直接比较。 6.4 奖励怎么配,往往比优化器更重要 真实系统几乎不会只用一个分数。以文字生成图片为例,可以同时有提示词一致性、OCR 正确率、审美、人体结构和安全性五类信号。最直接的做法是加权求和: $$R=w_{\text{prompt}}R_{\text{prompt}}+w_{\text{ocr}}R_{\text{ocr}}+w_{\text{aesthetic}}R_{\text{aesthetic}}-w_{\text{safety}}C_{\text{safety}}.$$ 但“都放进来”不等于问题解决了。不同奖励的量纲与波动范围可能差几个数量级:OCR 是 0 或 1,审美模型可能落在 1 到 10,安全惩罚可能只有极少数样本非零。若不先校准,权重看似相近,实际梯度却会被方差最大的那一项占满。稳妥做法是先在固定基座模型上采一批校准集,观察每项分布,再做裁剪、标准化或分位数映射;训练期间同时画出每个子奖励,不能只保存加权总分。 还要区分“软目标”和“硬约束”。图片里文字少一个字,通常是可连续改进的软目标;生成违法内容则不应靠一个可被其他高分抵消的负权重处理。硬约束更适合在采样前后用过滤器、拒绝策略或拉格朗日约束单独处理。否则模型可能发现:只要审美分足够高,安全惩罚也值得承受。 多目标之间还会发生真实冲突。把 OCR 权重拉高,模型可能把文字做得巨大而破坏构图;只奖励检测器置信度,可能得到边缘锐利但不自然的目标。应当用 Pareto 视角看结果:同一批 checkpoints 同时比较各子目标和人工偏好,选择没有被某一维明显支配的方案,而不是盯着一条总 reward 曲线挑最高点。 6.5 三个能快速定位实现错误的不变量 第一,rollout 保存的状态数必须比转移 log-prob 数多 1;若二者相等,通常漏了初始噪声或终点 latent。第二,当新旧模型参数完全相同时,逐步 ratio 应非常接近 1,approximate KL 应接近 0;否则多半是调度器、CFG 分支、时间步索引或 log-prob 归约维度不一致。第三,打乱 batch 顺序不应改变同一个样本的 reward、prompt、latents 与 log-prob 对齐关系。异步奖励最容易在这里悄悄错位:训练仍会运行,均值甚至会上升,但模型学到的是别人的分数。 调试时先用一个可以解析求解的低维过程,就像本文代码;确认梯度方向、ratio 和 baseline 都符合预期,再接入大模型。直接在多卡图像训练里找符号错误,通常会把昂贵计算浪费在最便宜的 bug 上。 这些不变量也适合写进自动化测试:用固定随机种子的十几条轨迹做快速回归,每次升级 diffusers、调度器或混合精度设置后先跑它,再启动昂贵的正式训练。 07. 经典论文脉络 Training Diffusion Models with Reinforcement Learning / DDPO(arXiv:2305.13301)把扩散去噪正式改写为多步决策过程,给出基于 likelihood ratio 的 DDPO,并展示了不可微奖励下的微调能力。它奠定了“去噪轨迹就是策略轨迹”的主框架。 DPOK: Reinforcement Learning for Fine-tuning Text-to-Image Diffusion Models(arXiv:2305.16381)把在线策略梯度与 KL 正则结合,强调奖励提升和预训练分布保持之间的平衡。 DRaFT: Directly Fine-tuning Diffusion Models on Differentiable Rewards(arXiv:2309.17400)转向可微奖励,通过采样链直接反传,并研究 full、截断和低方差版本,形成与 DDPO 互补的路线。 Reward-Directed Conditional Diffusion(arXiv:2307.07055)从 reward-conditioned / reward-directed 角度研究如何把黑盒目标引入扩散生成,说明“训练一个条件生成器”和“在线策略优化”是相关但不同的解法。 Diffusion Models for Reinforcement Learning: A Survey(arXiv:2311.01223)覆盖的是更宽的交叉方向:既包括用 RL 对齐扩散生成,也包括把扩散模型当作策略或规划器。阅读时要先辨认“RL 优化 diffusion”还是“diffusion 服务 RL”。 论文名称很像时,先看梯度从哪里来:若公式里是 $R\nabla\log p_\theta$,属于 score-function 策略梯度;若是 $\nabla_{x_0}R$ 沿采样链反传,则属于可微奖励路线;若训练数据是静态偏好对,通常又是离线偏好优化问题。 08. 常见误解 误解 1:模型预测的噪声就是 RL 动作 噪声预测是策略的参数化中间量。DDPO 的动作通常记为实际采样出的下一 latent $x_{t-1}$,策略概率才是 $p_\theta(x_{t-1}\mid x_t,c)$。只有这样,保存动作、重算 log-prob 和构造新旧比率才是一致的。 误解 2:奖励在最后,所以只有最后一步能更新 轨迹 log-prob 是逐步 log-prob 的和。终点奖励乘的是整条和,因此每个去噪转移都有 score-function 梯度。最后一步可能贡献较大,但不是唯一有梯度的一步。 误解 3:确定性 DDIM 也能原样计算 DDPO log-prob $\eta=0$ 时 $\sigma_t=0$,普通高斯 log-prob 不再成立。要么使用非零随机性,要么换成适合确定性路径的可微反传或其他估计器,不能简单给分母加一个 epsilon 就宣称理论等价。 误解 4:把所有逐步比率相乘,才是最“严格”的 PPO 轨迹比率在长链上方差巨大。DDPO 实践通常按去噪转移计算和裁剪 ratio;这是一种稳定的 surrogate,而不是把一个 50 步连乘数硬塞进普通 PPO 公式。 误解 5:同一奖励广播到每一步,所以每步学到的一样 优势相同,不等于 score function 相同。噪声方差、预测残差、网络敏感度和调度器权重都随时间步变化,逐步梯度自然不同。 误解 6:奖励不可微就无法优化扩散模型 DDPO 只需要采样结果和标量奖励,不需要奖励梯度。不可微带来的问题是估计方差和样本成本,而不是“没有任何梯度”。 误解 7:训练 reward 上升就代表模型全面变好 模型可能只学会利用 reward 的盲点。必须同时看独立评估、人工盲评、多样性、prompt 遵循和安全指标;最好用不同架构的 holdout reward 检查泛化。 09. 动手验证:四个改动看懂算法 文末代码只依赖 NumPy,建议依次改四个参数,每次固定随机种子并比较输出。 实验 A:把 steps 从 8 改为 16 预期状态 shape 从 (4096, 9) 变成 (4096, 17),log-prob shape 从 (4096, 8) 变成 (4096, 16)。更长轨迹会增加梯度求和项;在不调整学习率与归一化时,训练波动通常更明显。 实验 B:把 noise_std 从 0.45 降到 0.10 探索减弱,采样更集中;但 score function 含 $1/\sigma^2$,数值会更尖锐。过小噪声并不等于更稳定,极限到 0 反而失去这套高斯策略梯度的基础。 实验 C:删掉优势中心化 把 advantage = (reward - reward.mean()) / ... 改成直接使用 reward。gradient_sample_std_raw_centered 的第二个值会失去优势,训练曲线更抖。这个实验直观看到 baseline 改变方差而不改变目标最优点。 实验 D:增大新旧参数差 在 PPO 诊断部分增大 new_theta - old_theta。你会看到 ratio_min/max 更极端,clip_fraction 上升。工程上这对应学习率过大、inner epochs 过多或 KL 约束过弱。 进一步可把终点奖励改成离散黑盒:例如 reward = (abs(x0-target) < 0.25)。代码仍能训练,但由于奖励更稀疏,往往需要更大 batch、分层采样或更好的 baseline。这正是从玩具例子走向 OCR 成功率、目标检测命中和人工反馈时遇到的现实问题。 10. 延伸阅读 这篇位于“对齐与强化学习”路径的第三段: 策略梯度与 PPO 基础:理解 advantage、ratio、clipping 和 KL; 从 DPO 到 GRPO:去掉价值网络:理解组内相对优势; 本文 DiffusionRL:把 token 轨迹换成 latent 去噪轨迹; 视频生成中的强化学习与奖励模型:在去噪信用之外,再处理画质、运动与指令遵循的多目标奖励。 如果 DDIM 还不熟,先记住本文真正依赖的最小事实:它把 $x_t$ 和网络预测组合成下一步均值,并可用 $\eta$ 注入高斯随机性。从 DDIM 到高阶采样器 已展开采样公式、ODE/SDE 视角与高阶求解器。 最后用一句话收束: DiffusionRL 不是把一个 PPO 套在图片外面,而是把每次去噪转移变成可计概率的动作,再让终点奖励沿整条轨迹分配信用。 附录:完整代码 09 节用到的脚本全文如下(ddpo_minimal.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 ddpo_minimal.py """NumPy-only toy DDPO: treat iterative denoising as a stochastic policy. This is an algebra demo, not an image generator. A scalar latent follows a short Gaussian reverse process. One parameter shifts every denoising mean, and a terminal reward teaches the process to end near a target value. """ from dataclasses import dataclass import numpy as np @dataclass(frozen=True) class ToyConfig: steps: int = 8 batch_size: int = 4096 contraction: float = 0.72 action_scale: float = 0.18 noise_std: float = 0.45 target: float = 1.5 def gaussian_log_prob(value, mean, std): """Elementwise log N(value; mean, std^2).""" return -0.5 * ((value - mean) / std) ** 2 - np.log(std * np.sqrt(2.0 * np.pi)) def rollout(theta, cfg, rng): """Sample x_T -> ... -> x_0 and retain every transition. Array index 0 stores x_T; index cfg.steps stores x_0. The transition mean is mu_theta = contraction * x_t + action_scale * theta. """ latents = np.empty((cfg.batch_size, cfg.steps + 1), dtype=np.float64) means = np.empty((cfg.batch_size, cfg.steps), dtype=np.float64) log_probs = np.empty((cfg.batch_size, cfg.steps), dtype=np.float64) latents[:, 0] = rng.normal(size=cfg.batch_size) for j in range(cfg.steps): mean = cfg.contraction * latents[:, j] + cfg.action_scale * theta next_latent = mean + cfg.noise_std * rng.normal(size=cfg.batch_size) latents[:, j + 1] = next_latent means[:, j] = mean log_probs[:, j] = gaussian_log_prob(next_latent, mean, cfg.noise_std) terminal = latents[:, -1] rewards = -(terminal - cfg.target) ** 2 return latents, means, log_probs, rewards def score_function(latents, means, cfg): """d log pi_theta(x_{t-1}|x_t) / d theta for every transition.""" residual = latents[:, 1:] - means return residual * cfg.action_scale / (cfg.noise_std**2) def policy_gradient(rewards, per_step_scores): """REINFORCE with a batch baseline; terminal advantage is broadcast to T steps.""" centered = rewards - rewards.mean() advantages = centered / (rewards.std() + 1e-8) trajectory_scores = per_step_scores.sum(axis=1) gradient = np.mean(advantages * trajectory_scores) per_step_gradient = np.mean(advantages[:, None] * per_step_scores, axis=0) return gradient, per_step_gradient, advantages def evaluate(theta, cfg, seed): rng = np.random.default_rng(seed) latents, means, log_probs, rewards = rollout(theta, cfg, rng) return { "terminal_mean": float(latents[:, -1].mean()), "reward_mean": float(rewards.mean()), "latents": latents, "means": means, "log_probs": log_probs, "rewards": rewards, } def train(cfg, updates=40, learning_rate=0.08, seed=7): theta = 0.0 history = [] rng = np.random.default_rng(seed) for update in range(updates + 1): latents, means, _, rewards = rollout(theta, cfg, rng) scores = score_function(latents, means, cfg) gradient, per_step_gradient, advantages = policy_gradient(rewards, scores) history.append((update, theta, latents[:, -1].mean(), rewards.mean(), gradient)) if update < updates: theta += learning_rate * gradient return theta, np.asarray(history), per_step_gradient, advantages, scores, rewards def ppo_diagnostics(theta_old, theta_new, cfg, seed=123, clip_range=0.2): """Re-evaluate old transitions under a candidate new policy.""" rng = np.random.default_rng(seed) latents, old_means, old_log_probs, rewards = rollout(theta_old, cfg, rng) new_means = cfg.contraction * latents[:, :-1] + cfg.action_scale * theta_new new_log_probs = gaussian_log_prob(latents[:, 1:], new_means, cfg.noise_std) ratios = np.exp(new_log_probs - old_log_probs) clipped = np.abs(ratios - 1.0) > clip_range return rewards, ratios, clipped if __name__ == "__main__": np.set_printoptions(precision=6, suppress=True) cfg = ToyConfig() before = evaluate(theta=0.0, cfg=cfg, seed=2026) theta, history, per_step_gradient, advantages, scores, rewards = train(cfg) after = evaluate(theta=theta, cfg=cfg, seed=2026) raw_gradient_samples = rewards * scores.sum(axis=1) centered_gradient_samples = (rewards - rewards.mean()) * scores.sum(axis=1) ppo_rewards, ratios, clipped = ppo_diagnostics( theta_old=0.0, theta_new=0.5, cfg=cfg ) print("trajectory_shape=", before["latents"].shape) print("transition_log_prob_shape=", before["log_probs"].shape) print("before_terminal_mean=", f"{before['terminal_mean']:.4f}") print("before_reward_mean=", f"{before['reward_mean']:.4f}") print("after_theta=", f"{theta:.4f}") print("after_terminal_mean=", f"{after['terminal_mean']:.4f}") print("after_reward_mean=", f"{after['reward_mean']:.4f}") print("per_step_gradient=", np.round(per_step_gradient, 4)) print( "gradient_sample_std_raw_centered=", f"{raw_gradient_samples.std():.4f}", f"{centered_gradient_samples.std():.4f}", ) print( "ppo_ratio_mean_min_max=", f"{ratios.mean():.4f}", f"{ratios.min():.4f}", f"{ratios.max():.4f}", ) print("ppo_clip_fraction=", f"{clipped.mean():.4f}") print("history_first=", np.round(history[0, 1:4], 4)) print("history_last=", np.round(history[-1, 1:4], 4)) assert before["latents"].shape == (cfg.batch_size, cfg.steps + 1) assert before["log_probs"].shape == (cfg.batch_size, cfg.steps) assert after["reward_mean"] > before["reward_mean"] assert abs(after["terminal_mean"] - cfg.target) < abs(before["terminal_mean"] - cfg.target) assert np.all(np.isfinite(ratios)) make_figures.py """Generate the explanatory figures for the DiffusionRL article. Requires NumPy, Matplotlib and Pillow. The numerical chart reuses the exact toy process from ddpo_minimal.py so the article and figure cannot drift apart. """ from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np from PIL import Image, ImageDraw, ImageFont from ddpo_minimal import ToyConfig, train OUT = Path(__file__).resolve().parents[1] / "figures" OUT.mkdir(exist_ok=True) def _font(size, bold=False): candidates = ( "/System/Library/Fonts/PingFang.ttc", "/System/Library/Fonts/Hiragino Sans GB.ttc", "/usr/share/fonts/opentype/noto/NotoSansCJK-Bold.ttc" if bold else "/usr/share/fonts/opentype/noto/NotoSansCJK-Regular.ttc", ) for path in candidates: try: return ImageFont.truetype(path, size) except OSError: pass raise RuntimeError("请安装中文字体后再生成 DiffusionRL 示意图") def trajectory_mdp(): image = Image.new("RGB", (1800, 980), "#f8fafc") draw = ImageDraw.Draw(image) ink, blue, green, orange, muted = "#172033", "#2563eb", "#059669", "#ea580c", "#526078" draw.text((90, 55), "把多步去噪看成一条 MDP 轨迹", font=_font(70, True), fill=ink) draw.text((92, 150), "状态是当前 latent,动作是下一 latent 的采样,终点图像只得到一次奖励", font=_font(38), fill=muted) draw.text((900, 285), "策略 πθ = pθ(next | current, c)", font=_font(32), fill=blue, anchor="mm") xs = [170, 480, 790, 1100, 1410] labels = ["xT", "xT−1", "xT−2", "…", "x₀"] subtitles = ["纯噪声", "较粗结构", "结构成形", "继续去噪", "最终样本"] for i, (x, label, subtitle) in enumerate(zip(xs, labels, subtitles)): color = blue if i < len(xs) - 1 else green draw.rounded_rectangle((x - 105, 360, x + 105, 540), radius=36, fill="#ffffff", outline=color, width=6) draw.text((x, 400), label, font=_font(46, True), fill=color, anchor="mm") draw.text((x, 485), subtitle, font=_font(28), fill=muted, anchor="mm") if i < len(xs) - 1: draw.line((x + 110, 450, xs[i + 1] - 120, 450), fill=ink, width=5) draw.polygon( [(xs[i + 1] - 120, 450), (xs[i + 1] - 148, 433), (xs[i + 1] - 148, 467)], fill=ink, ) draw.line((1515, 450, 1650, 450), fill=orange, width=5) draw.polygon([(1650, 450), (1622, 433), (1622, 467)], fill=orange) draw.rounded_rectangle((1510, 630, 1735, 820), radius=32, fill="#fff7ed", outline=orange, width=5) draw.text((1622, 675), "R(x₀,c)", font=_font(44, True), fill=orange, anchor="mm") draw.text((1622, 750), "终点奖励", font=_font(31), fill=muted, anchor="mm") draw.line((1620, 545, 1620, 625), fill=orange, width=5) draw.polygon([(1620, 625), (1603, 597), (1637, 597)], fill=orange) draw.rounded_rectangle((90, 650, 1380, 850), radius=30, fill="#eef4ff") draw.text((130, 685), "同一个优势 A 被广播到每一步:", font=_font(39, True), fill=ink) draw.text((130, 750), "A · [grad log pθ(xT−1 | xT) + … + grad log pθ(x₀ | x₁)]", font=_font(39), fill=blue) draw.text((130, 805), "奖励只在终点出现,但整条轨迹的 log-prob 都参与更新。", font=_font(31), fill=muted) image.save(OUT / "diffusion_rl_mdp.png", optimize=True) def cover(): image = Image.new("RGB", (1280, 720), "#081225") draw = ImageDraw.Draw(image) # A denoising trajectory that gradually changes from noise-like dots into a # clean latent. Keeping this code-native makes the cover deterministic. rng = np.random.default_rng(7) stages = [190, 400, 610, 820, 1030] palette = ["#60a5fa", "#38bdf8", "#2dd4bf", "#34d399", "#f59e0b"] for i, x in enumerate(stages): radius = 62 draw.ellipse((x - radius, 258 - radius, x + radius, 258 + radius), outline=palette[i], width=5) spread = 47 - i * 7 for _ in range(38 - i * 4): px = x + int(rng.normal(0, spread)) py = 258 + int(rng.normal(0, spread)) dot = 2 + i draw.ellipse((px - dot, py - dot, px + dot, py + dot), fill=palette[i]) if i < len(stages) - 1: draw.line((x + 72, 258, stages[i + 1] - 75, 258), fill="#94a3b8", width=4) draw.polygon( [(stages[i + 1] - 75, 258), (stages[i + 1] - 96, 246), (stages[i + 1] - 96, 270)], fill="#94a3b8", ) draw.rounded_rectangle((1068, 190, 1222, 326), radius=26, fill="#431407", outline="#f59e0b", width=4) draw.text((1145, 258), "+R", font=_font(48, True), fill="#fbbf24", anchor="mm") draw.text((92, 415), "把 RL 用到扩散模型上", font=_font(64, True), fill="#f8fafc") draw.text((96, 507), "DiffusionRL · DDPO · 去噪轨迹 · 信用分配", font=_font(34), fill="#93c5fd") draw.rounded_rectangle((94, 603, 426, 659), radius=25, fill="#1d4ed8") draw.text((260, 631), "AIGC 基本功", font=_font(28, True), fill="#ffffff", anchor="mm") image.save(OUT / "wechat_cover_diffusion_rl.png", optimize=True) def training_and_credit(): cfg = ToyConfig() _, history, per_step_gradient, _, _, _ = train(cfg) updates = history[:, 0] fig, axes = plt.subplots(1, 2, figsize=(11, 4.3)) ax = axes[0] ax.plot(updates, history[:, 3], color="#2563eb", linewidth=2.5, label="mean terminal reward") ax.set_xlabel("policy updates") ax.set_ylabel("reward (higher is better)") ax.grid(alpha=0.2) ax2 = ax.twinx() ax2.plot(updates, history[:, 2], color="#059669", linewidth=2.2, label="mean x0") ax2.axhline(cfg.target, color="#059669", linestyle="--", alpha=0.45, label="target") ax2.set_ylabel("terminal latent mean") lines = ax.get_lines() + ax2.get_lines() ax.legend(lines, [line.get_label() for line in lines], frameon=False, loc="center right") ax.set_title("Toy DDPO learns from terminal reward") ax = axes[1] steps = np.arange(cfg.steps) ax.bar(steps, per_step_gradient, color=plt.cm.Blues(np.linspace(0.45, 0.95, cfg.steps))) ax.set_xticks(steps, [f"{cfg.steps-i}→{cfg.steps-i-1}" for i in steps], rotation=40) ax.set_xlabel("denoising transition") ax.set_ylabel("estimated gradient contribution") ax.grid(axis="y", alpha=0.2) ax.set_title("Same reward, different credit per step") fig.suptitle("Terminal reward becomes a trajectory-level policy gradient", fontsize=15, weight="bold") fig.tight_layout() fig.savefig(OUT / "diffusion_rl_training_credit.png", dpi=190) plt.close(fig) if __name__ == "__main__": trajectory_mdp() training_and_credit() cover() for name in ( "diffusion_rl_mdp.png", "diffusion_rl_training_credit.png", "wechat_cover_diffusion_rl.png", ): print(OUT / name) 更多 AIGC 论文解读,关注微信公众号「人工智能炼丹君」 每日更新 · 论文精选 · 深度解读 · 技术脉络 微信搜索 人工智能炼丹君 或扫描下方二维码关注
2026年09月17日
7 阅读
0 评论
0 点赞
2026-09-15
AIGC 基本功|视频 DiT 里的 3D RoPE 与分辨率外推-3D-RoPE
视频 DiT 里的 3D RoPE 与分辨率外推 所属方向:注意力与核心零件 | 难度:进阶 | 前置知识:RoPE 的原理与实现 关键词:3D RoPE、时空位置编码、视频 DiT、分辨率外推、逐轴位置插值、帧数外推 01. 为什么需要它 先看一对位置:视频格点 A=(0,0,0),B=(0,1,0)。B 在 A 的正下方一格。将每帧按行展开,再按时间串接,普通 1D RoPE 看到的序列位置是 m=tHW+hW+w。当每帧宽 W=4 时,A 到 B 的序列距离是 4;仅把宽改为 W=8,距离变成 8。两块仍上下相邻,旋转角却按 1D 的位置差变了。时间上的相邻帧也会受 HW 改变影响。固定画幅下,这种展开次序可以用,模型也可能学到空间规律;问题是在画幅或时长改变时,序列距离不再稳定代表同一种时空关系。这里的 1D RoPE 是为了隔离坐标问题而构造的诊断基线,不是在声称所有早期视频模型都按这一公式编码位置。 3D RoPE 给 A、B 保留 (t,h,w),所以无论宽度是多少,两者都只在高轴相差一格。它解决的是展平编号把画幅尺寸混进位置差的问题;画质外推还涉及训练分辨率、物体尺度和注意力成本。后文用 CogVideoX 的实验 分开讨论位置表直接扩展与逐轴插值,不把“坐标可计算”误当成“视频一定生成得好”。 02. 最小可用理解 基础 RoPE 已解释二维旋转与相对位置性质,这里直接看视频新增的两步: 换坐标:展平版给每个 token 一个 m,相对位移是 Δm;三轴版保留 (t,h,w),相对位移是 (Δt,Δh,Δw)。图 1 的上下邻居在画幅变宽后,只有三轴版仍把它识别为“高轴相差一格”。 分特征与频率:一个注意力头的特征分为时间、高、宽三段,每段用自己的坐标和频率表生成位置角。三段随后仍在同一次注意力内积里汇总;“3D”说的是坐标轴数,不是把视频物体在欧氏空间旋转。 图 1:固定 A、B 的时空坐标,只改每帧格点宽度。展平位置差从 4 变 8,三轴位移始终是 (0,1,0);真实模型的内容特征也会随画幅变化。 具体比例由模型选择。CogVideoX 给时间 1/4、高 3/8、宽 3/8;64 维头对应 (16,24,24),不是所有 3D RoPE 的固定标准。按这个分配,高轴的坐标映射可以单独调整,而时间、宽轴保持原样。实际换模型还要核对每轴维度、坐标单位、频率基数、配对次序和文本 token 的处理,不能只照搬比例。 图 2:64 维注意力头的示例。每个轴的维度都成对构成旋转平面;三段共享一个 token,但不会互相混成单一序列编号。图由文末代码生成。 03. 数学推导 3.1 从视频格点到三组旋转 设视频 patch token 为 $z_{t,h,w}$。$t$ 是压缩后的时间格点,$h$ 和 $w$ 是 patch 后的高、宽格点;它们不是原始帧号或像素坐标。以下只讨论注意力打分前,如何给对应的 query/key 构造位置相位。 先写出基础版到底把什么送进旋转。固定一个按行优先展开的视频网格 $(T,H,W)$,位置编号 $m=tHW+hW+w$。任意两格 $p=(t,h,w)$、$r=(t',h',w')$ 的编号差为: $$m_r-m_p=(t'-t)HW+(h'-h)W+(w'-w)$$ 上式说明 1D 并非丢了时空信息:固定 $H,W$ 时,每个格点仍有唯一编号。但同样一个高轴位移会乘上当前的 $W$,同样一个时间位移会乘上当前的 $HW$。换画幅后,“向下 1 格”的位置相位因此变化;这只比较固定内容特征的位置项,不推断整个模型的注意力分数或画质。 令头维度 $d=d_t+d_h+d_w$,每段维度都必须为偶数。把 $q$ 和 $k$ 拆为 $q^{(t)},q^{(h)},q^{(w)}$ 与对应的三段 $k$。对轴 $a\in\{t,h,w\}$ 的第 $i$ 个二维旋转平面,选频率: $$\omega_{a,i}=\Theta^{-2i/d_a},\qquad i=0,1,\ldots,d_a/2-1$$ $\Theta$ 是频率基数,CogVideoX 的源码默认 10000;$d_a$ 是该轴分到的维度,而不是整个头维度。时间与空间段各自从最高频重新开始,不是把一条 1D 频率数组切成三段。沿用先修文章的旋转记号 $R$,用 $R_a(p_a)$ 表示轴 $a$ 上所有旋转平面的组合。完整位置操作是三个互不混合的块拼接: $$\widetilde q_{t,h,w}=\operatorname{diag}(R_t(t),R_h(h),R_w(w))q_{t,h,w}$$ key 用同一个位置规则得到 $\widetilde k$。分轴的是位置相位的构造;注意力仍对所有可见视频 token 计算一次打分。 可以把区别压到一行:1D 的相对角是 Δm × 频率,其中 Δm=Δt·HW+Δh·W+Δw;3D 的三组相对角分别是 Δt × 时间频率、Δh × 高频率、Δw × 宽频率。三组贡献在同一次内积里相加。这里的 Δ 都是 token 格点差,不是秒数或像素差。 3.2 三个位移怎样汇入一次打分 基础 RoPE 已推过两个旋转角相减的恒等式。这里新增的是:每个轴只用自己的坐标差,整头把三段贡献相加: $$s(p,r)=\sum_{a\in\{t,h,w\}}\langle q^{(a)},R_a(r_a-p_a)k^{(a)}\rangle$$ $p=(t,h,w)$ 和 $r=(t',h',w')$ 是两个格点。只有高差为 1 时,时间段与宽段的位置差都是零,高段单独改变打分;换宽度不会把这一高差改写成另一个数。公式描述固定 $q,k$ 时的位置作用,文末代码用上下邻居和下一帧两种位移分别检查它。 再看一个跨帧例子:某个物体的投影从 (0,5,5) 到 (1,5,6),位置差是 (1,0,1);时间段和宽段的相位会变,高段不会。但静止物体遇到相机运动也可能产生完全相同的屏幕投影轨迹。对这两个例子,3D RoPE 的位置部分完全一样;它没有标记“这是同一物体”,也没有记录相机姿态或两个时间格点之间的秒数。模型仍可从图像内容、条件和后续层学习运动关系,不能把三轴相位本身当作光流或物理速度。 式子还有一个可检查的边界:三段打分相加,位置操作本身没有 Δt×Δw 这样的跨轴乘积项。这不妨碍注意力通过 $q,k$ 的内容组合时空线索;它只说明“分别编码三个轴”并不等于“直接编码一条物体轨迹”。 3.3 直接外推与逐轴插值到底改了什么 假设训练高轴有 $H_0$ 个格点,推理变成 $H_1$ 个格点,零起点时两者的最大坐标分别是 $L_0=H_0-1$、$L_1=H_1-1$。直接外推使用新坐标 $h'=h$,相邻高 token 仍相差 1,角差仍为 $\omega_{h,i}$,但远端坐标超过旧范围。若要让两个端点严格对齐,逐轴位置插值应使用 $h'=h/s_h$、$s_h=L_1/L_0=(H_1-1)/(H_0-1)$;相邻角差变为 $\omega_{h,i}/s_h$,最远坐标正好被压回 $H_0-1$。常见的 $H_1/H_0$ 是大网格下的近似,网格很短时两者不可混写。宽与时间可以分别选择 $s_w,s_t$: $$\phi_{a,i}(p_a)=\frac{p_a}{s_a}\omega_{a,i},\qquad s_t,s_h,s_w>0$$ 这条式子说明“分辨率外推”不是只能一起拉伸三个轴。若只想演示高轴坐标压缩两倍,可以选 $(s_t,s_h,s_w)=(1,2,1)$;若要复现某条真实管线,则应按它的端点、crop 坐标和取样约定计算,不能只看形状比。保持局部细节对应相邻格点的相位差尽量不变;覆盖新范围对应最远端相位别离开已训练的分布。单一全频率插值无法同时严格满足两者。图 3 用 scale=2 隔离这个机制:高度插值后高频振荡慢了一半,低频在这 32 个格点上原本就几乎不动;它不是某个 CogVideoX 推理配置的逐参数复刻。读图时也不要把 cos 的偶然重合当成“两个远处位置相同”。 图 3:高维度 24、基数 10000 的示例。左边用原坐标,右边只把高坐标除以 2;时间与宽并未压缩。图中画的是 cos 相位而非生成质量。 3.4 维度预算不是按视频长宽直接分蛋糕 为什么 CogVideoX 把时间分到 1/4、两个空间轴各分到 3/8,而不是三等分?论文给出了这个比例与三轴独立 RoPE 的实现,却没有证明它对所有数据和画幅最优。可以从编码需求理解:空间中每一帧都包含大量局部邻接与大范围结构,两个空间轴合计拿到较多旋转平面;压缩后的时间格点通常少于空间格点,但仍需能分辨前后帧。这个解释是设计动机的推断,不能代替原论文消融,更不能据此宣布时间维度永远够用。 64 维示例中,时间 16 维对应 8 个旋转平面,高与宽的 24 维各对应 12 个。最高频平面的角增量都是 1 弧度,因为 $i=0$ 时 $\omega=1$,所以“空间分得多”并不是说高、宽的每一步比时间转得快。区别落在可用的频率层级数量和最慢频率:时间最后一组在坐标 15 的角约 0.004743,高度最后一组约 0.003232。两组数字都很小,它们在短网格上几乎不提供相邻区分,却可能给更长距离提供不同的相位尺度。把更多维度分给某轴,得到的是更丰富的频谱,不是对该轴长度的直接承诺。 还要区分物理帧数与时间 RoPE 格点。假设 VAE 把若干帧压缩到一个 latent 时间格点,视频从 49 帧变成 97 帧,RoPE 的时间长度未必简单从 49 变到 97;VAE 的首帧保留或补帧规则甚至会让两个公式对不上。头维度预算必须在真正进入注意力的 token 网格上讨论。读实现时先打印 patchify 输出和 (T,H,W),再看时间频率有没有被拉伸。 帧率变化不等于时长外推。 同样是把原始帧数翻倍,可能是同一段动作以两倍帧率采样,也可能是原帧率下播放两倍时间。若 VAE 与 patchify 最终都把它们变成两倍的时间格点,整数时间 RoPE 看到的范围一样,却不能仅凭位置角判断“相邻格点代表多少秒”。前者主要考察更密的动作采样,后者还考察更长的主体与场景记忆;测试时间缩放前应固定并记录帧率、实际秒数、latent 时间压缩和条件时间戳。这是视频轴特有的语义问题,空间高宽轴没有对应的播放时钟。 3.5 NTK 变基数为什么不是“插值”的同义词 另一类策略不缩坐标,而是逐轴更改频率基数 $\Theta_a$。若 $\Theta_a$ 变大,$i=0$ 的最高频 $\omega_{a,0}=1$ 完全不变,较低频平面转得更慢;于是近邻分辨率较容易保留,远距离相位覆盖可以加宽。它与 $p_a/s_a$ 把所有频率一起压低不同。NTK-aware、YaRN 等方法在语言长上下文中发展,迁移到视频时还要决定给哪个轴、哪些频段用、是否微调和怎样处理训练画幅。仅把 rope_theta 调大,不能替代逐轴分辨率训练或保证画质。 “频率变慢”也不等于“完全不损局部”。注意力打分汇总多个平面,改低频仍会改变其对近处和远处匹配的贡献。任何外推方法都必须用目标模型、目标画幅与目标时长做消融。 04. 代码实现 code/rope_3d_minimal.py 只依赖 NumPy。核心代码先给时间、高、宽各生成角度表,再广播到 (T,H,W),沿最后一维拼接成每个 token 的角度;rotate 把相邻两维作为一个二维平面旋转。它保留 float64 以让代数测试更容易观察,不模拟工业模型的 BF16 或 CUDA 性能。 dt, dh, dw = axis_dims(head_dim) at = axis_angles(np.arange(frames), dt, scale=scales[0])[:, None, None, :] ah = axis_angles(np.arange(height), dh, scale=scales[1])[None, :, None, :] aw = axis_angles(np.arange(width), dw, scale=scales[2])[None, None, :, :] angles = np.concatenate( [np.broadcast_to(a, (frames, height, width, d // 2)) for a, d in zip((at, ah, aw), (dt, dh, dw))], axis=-1 ).reshape(frames * height * width, head_dim // 2) cos = np.repeat(np.cos(angles), 2, axis=-1) sin = np.repeat(np.sin(angles), 2, axis=-1) 这里 angles 的最后一维只有 d/2 个旋转平面;repeat(...,2) 才让 cos/sin 与 d 维特征逐元素相乘。把代码复制到本地运行 python3 rope_3d_minimal.py,本次真实输出如下: axis_dims= (16, 24, 24) cache_shape= (24, 64) input_shape= (24, 64) rotated_shape= (24, 64) norm_error= 1.78e-15 same_offset_scores= 2.208374 2.208374 difference= 1.33e-15 vertical_neighbor_1d_offsets= 4 8 vertical_neighbor_1d_scores= 1.678217 3.749697 vertical_neighbor_3d_offsets= (0, 1, 0) (0, 1, 0) vertical_neighbor_3d_scores= 0.419453 0.419453 height_phase_at_15_first_last= [15. 0.003232] height_PI_x2_first_last= [7.5 0.001616] time_phase_at_15_first_last= [15. 0.004743] token_count_13x30x45= 17550 token_count_13x60x90= 70200 dense_attention_ratio= 16.0 naive_score_matrix_fp16_gib_per_head= 0.574 9.179 四条 vertical_neighbor_* 输出固定同一组随机 $q,k$ 与同一对上下相邻 token,只切换展平时的 W=4/8。1D 的 Δm 从 4 变 8,内积分数从 1.678217 变 3.749697;3D 的逐轴距离始终 (0,1,0),位置项分数都为 0.419453。1D 与 3D 的绝对分数本身不该互比:两套频率配置不同;这个小实验只比较各自规则在宽度变化前后是否稳定。模型实际改宽后 $q,k$ 也会变,所以它不是画质或整层注意力不变的证据。 cache_shape=(24,64) 来自 2×3×4=24 个视频格点。旋转前后范数误差接近浮点舍入,因为旋转矩阵保长度。高坐标 15 的最高频角为 15 弧度,插值后变成 7.5;高轴最慢一组角从 0.003232 变 0.001616。时间角仍是 15 和 0.004743,证明只改高轴没有动时间表。最慢高频与最慢时频数值不同,因为两轴拿到的维度数分别是 24 与 16,指数分母不同。 再看代码里一个表面上很普通的选择:scales 被明确写成 (time,height,width) 三元组,而不是一个全局的 scale=2。这个接口让实验可以只改一个轴,也强迫调用者记录每轴真实扩张比例。假如训练网格是 T0×H0×W0、目标网格是 T1×H1×W1,对格点数大于 1 的轴,端点对齐的起点应分别按 (T1-1)/(T0-1)、(H1-1)/(H0-1)、(W1-1)/(W0-1) 计算;只有一个格点的轴没有可缩放的坐标跨度,应单独保持在零点。但这些比例仍不能机械照搬:训练数据往往包含多个时长和长宽比,没有单一的 T0,H0,W0。若宽度原本就见过目标值,宽轴可以维持原坐标;若时间增长跨越训练分布,更慢的低频策略可能更值得先测试。参数名称应表达“坐标映射”,而不是含糊写成“画质增强”。 norm_error、same_offset_scores 和宽度变化前后的 3D 分数在脚本里都有断言;若配对、广播或相对位置性质被改坏,脚本会直接失败,而不是只打印一串看似正常的数字。断言检查的是代数不变量,仍不能判断位置映射对视频是否有用,真实效果还要在固定权重的 DiT 上测。 请注意:示例里的 13×30×45 是格点形状。不能仅凭像素分辨率就猜它。真实模型要逐层根据 VAE 的时空压缩、首帧特殊处理、patch size 和 padding 算出 (T,H,W);错一个维度,位置表和 token 顺序都会错位。 05. 工业级实现对照 以 2026-09 核对的 CogVideo SAT 源码 Rotary3DPositionEmbeddingMixin 为例,类中 dim_t = hidden_size_head // 4、dim_h = dim_w = hidden_size_head // 8 * 3,每轴各自建立频率与格点,再广播、拼接、取 sin/cos。它在注意力函数中只旋转视频 query/key,先切出 text_length 个文本 token 原样放回;可选 rot_v 才旋转 value,默认并非必要。rope_T/H/W 从预计算的三维表里截取当前视频尺寸,排平为 (T*H*W,D)。这与最小代码的数组顺序一致,也便于逐轴检查实际相位。 一个容易踩的坑是:这个 SAT 类的构造函数虽然接收 height_interpolation、width_interpolation 和 time_interpolation 参数,在本次读取的 main 文件中,角度表仍直接由 torch.arange 计算,这些参数没有参与坐标缩放。因此不能因为配置里写了 height_interpolation=1.875,就声称该 SAT 旋转实现真的进行了高度插值。文章讨论的“逐轴插值”是可选的设计实验,不是这一段 CogVideo SAT 代码的实际行为。作者原论文采用扩表取前段,恰好与这里的原坐标逻辑相符。 Diffusers 的 get_3d_rotary_pos_embed 也有时间 1/4、空间各 3/8 的分配,但它提供 grid_type="linspace" 和 "slice" 两种格点构造。linspace 可把某个目标空间范围映射到当前高宽格点;slice 在最大表上用原索引再取当前尺寸。二者的差别首先是坐标选取,不是旋转公式变化。该函数在本次核对中返回 (cos,sin),供注意力处理器作用在 query/key 上。具体使用哪个 grid_type、裁剪坐标是什么,应沿管线调用栈核对,不能只看函数定义就给整个模型下结论。 写工业调用栈时建议按四个具体问题逐层追:输入是像素帧还是 VAE latent;patchify 后一条样本的 (T,H,W) 分别是多少;三轴坐标从零开始、按最大表切片还是重新映射到训练范围;cos/sin 最后究竟旋转的是视频 query/key、文本 token 是否另有位置处理。单独找到 get_3d_rotary_pos_embed 的名字,只能证明功能存在,不能证明这一模型版本已经启用它。CogVideoX 系列可能根据配置选择旋转位置或学习位置,不同 checkpoint 更不能一概而论。 维度排列也值得实际打印。本文的 rotate 假设相邻两维组成 (实部, 虚部),因此角表用 repeat_interleave(2) 形式复写到两维。若上游先把所有实部放前半、虚部放后半,旋转公式仍可成立,但 rotate_half 与 cos/sin 的广播顺序必须一起换。仅把缓存从另一仓库复制过来,形状同为 [N,D] 也可能在不报错的情况下给完全不同的旋转平面。检查一个位置 1、一个二维平面 (1,0) 的旋转结果应近似 (cos 1,sin 1),比只核对 shape 更可靠。 生产实现还会处理变长视频、文本与视频拼接、packing、缓存、设备和精度。最小脚本只面向一个规则视频网格;真实 batch 如果把不同视频压成一条序列,必须保存每段的起点与独立 (T,H,W),否则第二个视频的时间会错误地延续第一个视频的末帧。cos/sin 表可以复用,但格点顺序和真实 token 排列必须逐项对齐。很多“换分辨率立即崩”其实是这种形状或坐标实现错误,而非外推理论失败。 06. 代价与边界 纯 3D RoPE 本身不引入新的可训练位置参数,换形状时也可以按轴计算新表;具体模型仍可能另外叠加可学习的位置向量。这个性质让多分辨率训练与不同长度输入更自然,但没有把注意力从二次复杂度变成线性。高、宽各两倍时,视频 token 四倍、稠密注意力项十六倍:示例从 13×30×45=17,550 个 token 增至 13×60×90=70,200 个 token,若天真地为单个头物化一张 FP16 分数矩阵,理论容量会从约 0.574 GiB 涨到 9.179 GiB,尚未计梯度、softmax 和其他头。FlashAttention 不必把整张矩阵常驻显存,所以这不是实际峰值预测;它仍准确反映稠密注意力的二次计算量。把空间 RoPE 频率调得再聪明也不会消除这笔账,仍需 VAE 压缩、稀疏或局部注意力、分块计算等手段。 外推的另一重代价是位置区分与训练分布之间的冲突。直接延长表,邻近 token 的相位关系不变,却把远端坐标送到没见过的位置;全部插值,把远端压回旧尺度,却弱化每一步的相位变化。轴越短、分到的旋转平面越少,可用的频率跨度通常越窄。CogVideoX 的时间 16 维并不意味着“时间必然先坏”或“空间必然先坏”:结果还取决于训练长度、VAE 时间压缩、运动跨度与内容建模。 分辨率外推尤其要保住物理尺度解释。同一人物在原图占 10 个 latent 格点,放大画幅后占 20 个,坐标间隔与模型学过的视觉尺度同时变了;只插值位置并不能把纹理、构图与物体大小的分布一起搬过去。长视频类似:更多帧要求模型记住人物、场景和动作因果,位置相位不是记忆网络。若输出失败,先分辨是表索引越界、token 排列错位、局部重复、全局模糊还是跨帧身份漂移,再决定该调位置、数据、注意力或去噪。 相同 token 数也不意味着相同空间外推。 假设训练格点为 (T,H,W)=(8,32,32):目标 (8,64,32) 与 (8,32,64) 都把视频 token 翻倍,但前者只扩高轴,后者只扩宽轴。在 3D RoPE 中,它们分别改变高段或宽段的最远相位;在展平 1D RoPE 中,改宽还会重新标定所有行间与帧间的序列距离。实验可以保持 token 数、权重、prompt 和种子一致,只互换扩张轴,再分别检查竖向重复、横向重复与主体比例。若两个结果不同,先核对训练画幅分布和 VAE/patchify,才讨论是否来自轴频率。 排查失败时可以先做一个不动权重的三步实验。第一步固定 VAE latent 与 patch token,只换位置表构造,在同一坐标对上比较 q·k;若结果突然变化,检查表索引、轴顺序或维度配对。第二步固定画幅和视频内容,单独改变时间格点长度,看主体漂移是否只在跨越训练时长后出现。第三步固定时间格点,单独变高、宽并分别观察全局构图和局部纹理;若高宽都扩大导致局部重复,不能直接断言是 rope_theta 错,还要看是否出现 token 数暴涨引发的注意力截断、分块边界或 VAE tiling 伪影。把这些变量一次全改掉,会失去定位能力。 若比较位置插值与直接外推,评分也至少拆成两类:近邻细节,例如边缘、纹理与细小物体是否保持;远距结构,例如人物是否只有一个、左右场景是否连贯、同一主体跨帧是否一致。CogVideoX 附录里“一张模糊的大画面”与“多个清楚的小画面”的差异正说明:一个单一的美学分数可能掩盖局部与全局朝相反方向变化。画质指标之外还要记录 token 长度、时空网格与推理显存,否则方法之间可能使用了不同计算预算。 这套分轴设计也不表达“这是第 1 帧”这类绝对锚点——它主要编码相对位移。首帧图像条件、相机轨迹和剪辑时间戳需要通过额外条件或掩码表达。多图参考、回放状态与不同视频段拼接时,不同段若都从 (0,0,0) 重启,还需要明确段 ID 或相机等身份信息,免得同坐标误当同内容。 更远一步,镜头绕场景移动后再次拍到同一处墙面,这处墙面的世界位置没变,屏幕上的 (h,w) 却可能不同;反过来,相同 (h,w) 也可能指向另一块墙面。视频 3D RoPE 的“3D”是时间加二维屏幕格点,不是场景的三维坐标。面向相机可控的视频 DiT,SCoPE 将相机视线作为另一种位置线索加入注意力,并保留原有 RoPE;这是研究者针对屏幕坐标局限提出的扩展,不是普通 3D RoPE 自动拥有的几何能力。 07. 经典论文脉络 RoFormer、Position Interpolation 与 YaRN 的旋转及语言长上下文扩展已在先修文章讲过;对视频的问题是哪一轴、哪一频段需要缩放,语言模型结果不能直接当画质证据。 CogVideoX 将三维相对位置用于视频 DiT,并用多分辨率 Frame Pack 与渐进训练支持不同尺寸;论文附录展示的初始生成状态对比揭示了“局部清晰”与“全局成形”的矛盾。它是该模型上的定性证据,不是所有视频 RoPE 的通用胜负结论。 Diffusers 的 CogVideoX 实现 提供可核对的三轴维度分配、格点生成和注意力接口;论文、SAT 与 Diffusers 版本未必走同一坐标路径。 SCoPE 在视频 DiT 中保留时空格点 RoPE,另加相机视线线索以改善相机轨迹下的场景一致性;它讨论的是屏幕位置之外的几何信息,并非改写基础 RoPE 的旋转恒等式。 08. 常见误解 误解一:3D RoPE 就是把三个坐标加起来再做 1D RoPE。 若先算 $t+h+w$,(1,0,0) 与 (0,1,0) 就变成同一个位置值,时间变化和上下位移无法区分。正确做法是三段特征独立旋转,最后只在注意力内积里汇总贡献。 误解二:相对位置恒等式成立就证明任意分辨率都能生成。 恒等式对没见过的坐标同样成立;模型从未学过如何在新相位、更多 token 和新物体尺度上使用它。数学可计算性只保证不会因公式本身报错。 误解三:把所有坐标除以两倍没有副作用。 位置插值让新端点回到旧范围,却也让邻近格点的相位差减半,可能伤害局部细节。若目标只换高度,就不要无故压缩时间和宽度。 误解四:调大 rope_theta 等于 Position Interpolation。 更改基数主要影响较低频平面,最高频不动;位置除以缩放因子让所有频率一起慢下来。两者可能组合,作用机制不同。 误解五:CogVideoX 配置里有 height_interpolation 参数,就一定在使用插值。 本次读到的 SAT Rotary3DPositionEmbeddingMixin 接受该参数,却没有用它来计算 grid_h。要按真实执行路径确认;不同实现不能只凭参数名互相代替。 误解六:空间轴拿了 75% 维度,所以换分辨率导致的所有错误都来自空间 RoPE。 维度比例只是一种频率预算,无法排除注意力规模、训练分辨率、VAE 压缩和内容尺度的影响。看输出症状、做逐项消融,才有因果结论。 09. 动手验证 先运行 python3 rope_3d_minimal.py,看 vertical_neighbor_* 四行。把 lower=(0,1,0) 改成 (1,0,0),即从“下一行”改为“下一帧”:3D 位移由 (0,1,0) 改成 (1,0,0),1D 位移则由 W 改成 H×W。再单独改网格宽度,检查三轴版的下一帧相位是否仍只由时间坐标决定。 调用 rope_3d_cache(2,6,4,64,scales=(1,2,1)) 与 scales=(2,1,1),分别只压缩高、时间坐标;检查同一高坐标的高段相位和同一时间坐标的时间段相位。最后把网格 2×3×4 改成 2×6×4:缓存行数应从 24 变 48,列数仍为 64,只有高坐标的取值范围增加。 如果要评估生成模型的真实外推质量,不能只跑这段 NumPy 脚本。固定模型权重、prompt、种子与去噪设置,比较原训练画幅、直接扩表、逐轴插值和不同频段缩放;同时记录局部纹理、重复图案、全局构图、跨帧主体一致性与实际显存/延迟。只有这样才能分清“位置编码问题”与“高分辨率本来就需要新训练”的差异。 10. 延伸阅读 继续向下可以看视频 DiT 的稀疏注意力、3D VAE/patchify 和长视频状态记忆:位置只负责“在哪里”,这几项决定“看哪些 token、用什么尺度和怎样记住内容”。 本篇用到的具体外部依据以 CogVideoX 论文、CogVideo SAT 源码 与 Diffusers 位置编码源码 为准;它们分别支持论文中的训练取舍与两种工程实现,文中的 NumPy 输出仅证明所写公式和示例代码一致。 附录:完整代码 09 节用到的脚本全文如下(rope_3d_minimal.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 rope_3d_minimal.py """NumPy-only 3D RoPE and per-axis position extrapolation demo.""" import numpy as np def axis_dims(head_dim): """CogVideoX's time/height/width split for a head divisible by sixteen (all three axis sizes must be even).""" assert head_dim > 0 and head_dim % 16 == 0 dims = (head_dim // 4, head_dim * 3 // 8, head_dim * 3 // 8) assert all(d % 2 == 0 for d in dims) return dims def axis_angles(positions, dim, base=10000.0, scale=1.0): """Angles [positions, dim/2]; scale>1 compresses this axis only.""" assert dim % 2 == 0 and scale > 0 inv_freq = base ** (-np.arange(0, dim, 2, dtype=np.float64) / dim) return np.asarray(positions, dtype=np.float64)[:, None] * inv_freq[None, :] / scale def rope_3d_cache(frames, height, width, head_dim=64, scales=(1, 1, 1)): """Return interleaved cos/sin, one row per (t,h,w) token.""" assert frames > 0 and height > 0 and width > 0 assert len(scales) == 3 and all(scale > 0 for scale in scales) dt, dh, dw = axis_dims(head_dim) at = axis_angles(np.arange(frames), dt, scale=scales[0])[:, None, None, :] ah = axis_angles(np.arange(height), dh, scale=scales[1])[None, :, None, :] aw = axis_angles(np.arange(width), dw, scale=scales[2])[None, None, :, :] shape = (frames, height, width) angles = np.concatenate( [np.broadcast_to(a, shape + (d // 2,)) for a, d in zip((at, ah, aw), (dt, dh, dw))], axis=-1, ).reshape(frames * height * width, head_dim // 2) return np.repeat(np.cos(angles), 2, axis=-1), np.repeat(np.sin(angles), 2, axis=-1) def rotate(x, cos, sin): """Rotate adjacent dimension pairs; x and cache both [tokens, head_dim].""" paired = x.reshape(x.shape[0], -1, 2) quarter_turn = np.stack((-paired[..., 1], paired[..., 0]), axis=-1).reshape(x.shape) return x * cos + quarter_turn * sin def score_for_positions(q, k, pos_q, pos_k, scales=(1, 1, 1)): """A single dot product without building an attention matrix.""" dims = axis_dims(q.size) out_q, out_k = [], [] start = 0 for axis, dim in enumerate(dims): aq = axis_angles([pos_q[axis]], dim, scale=scales[axis]) ak = axis_angles([pos_k[axis]], dim, scale=scales[axis]) cq, sq = np.repeat(np.cos(aq), 2, axis=-1), np.repeat(np.sin(aq), 2, axis=-1) ck, sk = np.repeat(np.cos(ak), 2, axis=-1), np.repeat(np.sin(ak), 2, axis=-1) out_q.append(rotate(q[start:start + dim][None], cq, sq)[0]) out_k.append(rotate(k[start:start + dim][None], ck, sk)[0]) start += dim return np.dot(np.concatenate(out_q), np.concatenate(out_k)) def flatten_index(pos, height, width): """The row-major 1D index of a video token at (time, height, width).""" t, h, w = pos return t * height * width + h * width + w def score_for_flattened_positions(q, k, pos_q, pos_k, height, width): """Baseline: one 1D RoPE angle table for all head channels.""" mq = flatten_index(pos_q, height, width) mk = flatten_index(pos_k, height, width) aq = axis_angles([mq], q.size) ak = axis_angles([mk], k.size) cq, sq = np.repeat(np.cos(aq), 2, axis=-1), np.repeat(np.sin(aq), 2, axis=-1) ck, sk = np.repeat(np.cos(ak), 2, axis=-1), np.repeat(np.sin(ak), 2, axis=-1) return float(np.dot(rotate(q[None], cq, sq)[0], rotate(k[None], ck, sk)[0])) def naive_score_matrix_gib(token_count, bytes_per_element=2): """GiB for one materialized dense attention-score matrix of one head.""" return token_count * token_count * bytes_per_element / 2**30 if __name__ == "__main__": np.set_printoptions(precision=6, suppress=True) dims = axis_dims(64) cos, sin = rope_3d_cache(2, 3, 4, 64) rng = np.random.default_rng(7) tokens = rng.normal(size=cos.shape) rotated = rotate(tokens, cos, sin) norm_error = np.max(np.abs(np.linalg.norm(tokens, axis=1) - np.linalg.norm(rotated, axis=1))) print("axis_dims=", dims) print("cache_shape=", cos.shape, "input_shape=", tokens.shape) print("rotated_shape=", rotated.shape) print("norm_error=", f"{norm_error:.2e}") q, k = rng.normal(size=64), rng.normal(size=64) a = score_for_positions(q, k, (0, 1, 2), (1, 2, 3)) b = score_for_positions(q, k, (5, 6, 7), (6, 7, 8)) print("same_offset_scores=", f"{a:.6f}", f"{b:.6f}", "difference=", f"{abs(a-b):.2e}") upper, lower = (0, 0, 0), (0, 1, 0) flat4 = flatten_index(lower, 2, 4) - flatten_index(upper, 2, 4) flat8 = flatten_index(lower, 2, 8) - flatten_index(upper, 2, 8) one_d4 = score_for_flattened_positions(q, k, upper, lower, 2, 4) one_d8 = score_for_flattened_positions(q, k, upper, lower, 2, 8) three_d4 = score_for_positions(q, k, upper, lower) three_d8 = score_for_positions(q, k, upper, lower) three_d_offset = tuple(b - a for a, b in zip(upper, lower)) print("vertical_neighbor_1d_offsets=", flat4, flat8) print("vertical_neighbor_1d_scores=", f"{one_d4:.6f}", f"{one_d8:.6f}") print("vertical_neighbor_3d_offsets=", three_d_offset, three_d_offset) print("vertical_neighbor_3d_scores=", f"{three_d4:.6f}", f"{three_d8:.6f}") # These are algebraic invariants, not video-quality tests. Keeping them as # assertions turns the example into a small regression test for pairing, # broadcasting and relative-position behavior. assert norm_error < 1e-12 assert abs(a - b) < 1e-12 assert abs(three_d4 - three_d8) < 1e-12 assert not np.isclose(one_d4, one_d8) raw = axis_angles([15], dims[1])[0] pi = axis_angles([15], dims[1], scale=2)[0] print("height_phase_at_15_first_last=", np.round(raw[[0, -1]], 6)) print("height_PI_x2_first_last=", np.round(pi[[0, -1]], 6)) print("time_phase_at_15_first_last=", np.round(axis_angles([15], dims[0])[0][[0, -1]], 6)) small_tokens = 13 * 30 * 45 large_tokens = 13 * 60 * 90 print("token_count_13x30x45=", small_tokens) print("token_count_13x60x90=", large_tokens) print("dense_attention_ratio=", (large_tokens / small_tokens) ** 2) print( "naive_score_matrix_fp16_gib_per_head=", f"{naive_score_matrix_gib(small_tokens):.3f}", f"{naive_score_matrix_gib(large_tokens):.3f}", ) make_figures.py """Generate 3D RoPE figures. Requires NumPy, Matplotlib and Pillow.""" from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np from PIL import Image, ImageDraw, ImageFont from rope_3d_minimal import axis_angles, axis_dims OUT = Path(__file__).resolve().parents[1] / "figures" OUT.mkdir(exist_ok=True) def _chinese_font(size): for path in ( "/System/Library/Fonts/Hiragino Sans GB.ttc", "/System/Library/Fonts/PingFang.ttc", "/usr/share/fonts/truetype/noto/NotoSansCJK-Regular.ttc", ): try: return ImageFont.truetype(path, size) except OSError: pass raise RuntimeError("请安装中文字体后再生成 1D/3D 对照图") def one_d_vs_three_d(): """Show the same vertical neighbor under two video widths.""" image = Image.new("RGB", (1800, 950), "#f7f9fe") draw = ImageDraw.Draw(image) ink, muted = "#1e293b", "#475569" draw.text((80, 50), "同一对上下相邻的 token,换宽度会怎样?", font=_chinese_font(68), fill=ink) draw.text((82, 150), "A=(0,0,0) → B=(0,1,0);只改变视频格点宽度 W", font=_chinese_font(42), fill=muted) cards = ( (70, 260, 865, 790, "#e8efff", "#1d4ed8", "1D RoPE:一个序列距离", "W=4 → Δm=4", "W=8 → Δm=8", "上下相邻仍是 1 格,但旋转相位变了"), (935, 260, 1730, 790, "#e3f7f2", "#047857", "3D RoPE:三个轴分别计数", "W=4 → (0,1,0)", "W=8 → (0,1,0)", "高轴始终转 1 格,时间和宽轴不动"), ) for x0, y0, x1, y1, bg, accent, title, first, second, note in cards: draw.rounded_rectangle((x0, y0, x1, y1), radius=38, fill=bg) draw.rectangle((x0 + 38, y0 + 48, x0 + 50, y0 + 130), fill=accent) draw.text((x0 + 77, y0 + 51), title, font=_chinese_font(51), fill=ink) draw.text((x0 + 75, y0 + 210), first, font=_chinese_font(68), fill=accent) draw.text((x0 + 75, y0 + 322), second, font=_chinese_font(68), fill=accent) draw.text((x0 + 75, y0 + 435), note, font=_chinese_font(35), fill=muted) draw.text((80, 844), "旋转和相对位置性质相同;变化的是坐标系与每轴的特征维度。", font=_chinese_font(42), fill=ink) image.save(OUT / "one_d_vs_three_d.png", optimize=True) def axis_budget(): dims = axis_dims(64) fig, ax = plt.subplots(figsize=(9, 2.8)) colors = ["#1d4ed8", "#0d9488", "#ea580c"] start = 0 for name, dim, color in zip(("Time 25%", "Height 37.5%", "Width 37.5%"), dims, colors): ax.barh([0], [dim], left=start, color=color, height=0.55) ax.text(start + dim / 2, 0, f"{name}\n{dim} dims", ha="center", va="center", color="white", weight="bold") start += dim ax.set_xlim(0, 64) ax.set_ylim(-0.55, 0.55) ax.set_xlabel("Head channels (64 total)") ax.set_yticks([]) ax.set_title("CogVideoX 3D RoPE: independent rotary planes per axis") fig.tight_layout() fig.savefig(OUT / "axis_budget.png", dpi=180) plt.close(fig) def phase_tradeoff(): dim_h = axis_dims(64)[1] pos = np.arange(32) original = axis_angles(pos, dim_h) compressed = axis_angles(pos, dim_h, scale=2) fig, axes = plt.subplots(1, 2, figsize=(9, 3.5), sharey=True) for ax, values, title in zip(axes, (original, compressed), ("Raw coordinates", "Height PI: positions / 2")): ax.plot(pos, np.cos(values[:, 0]), color="#ea580c", linewidth=2, label="highest frequency") ax.plot(pos, np.cos(values[:, -1]), color="#1d4ed8", linewidth=2, label="lowest frequency") ax.set_title(title) ax.set_xlabel("Height-token index") ax.grid(alpha=0.18) axes[0].set_ylabel("cos(angle)") axes[1].legend(frameon=False, loc="lower right") fig.suptitle("Interpolation halves phase change on this axis at every frequency") fig.tight_layout() fig.savefig(OUT / "phase_tradeoff.png", dpi=180) plt.close(fig) if __name__ == "__main__": one_d_vs_three_d() axis_budget() phase_tradeoff() for name in ("one_d_vs_three_d.png", "axis_budget.png", "phase_tradeoff.png"): print(OUT / name) 更多 AIGC 论文解读,关注微信公众号「人工智能炼丹君」 每日更新 · 论文精选 · 深度解读 · 技术脉络 微信搜索 人工智能炼丹君 或扫描下方二维码关注
2026年09月15日
7 阅读
0 评论
0 点赞
2026-09-13
AIGC 基本功|从 DPO 到 GRPO:去掉价值网络-GRPO
从 DPO 到 GRPO:去掉价值网络 所属方向:对齐与强化学习 | 难度:进阶 | 前置知识:PPO 关键词:DPO、GRPO、偏好优化、组内归一化、免价值网络、参考模型 01. 为什么需要它 把一个 7B 语言模型接进 PPO-RLHF,训练栈很快会变成四个模型:负责生成的 actor、预测每个 token 未来回报的 critic、限制策略漂移的 reference model,以及把整条回答变成分数的 reward model。真正更新参数的是 actor 和 critic,后两者虽然通常冻结,前向和显存仍要付钱。 critic 是最难解释的一笔成本。语言模型往往只在回答结束时拿到一个总分,但 value head 必须给回答中每个 token 估计未来回报。它不但要和 actor 同量级地跑前向、反向,还要从稀疏的终局信号中学习密集价值。DeepSeekMath 指出,LLM 的 value model 通常与 policy model 规模相当,这正是 GRPO 要删除的部件。 以全参数 AdamW 的粗略账本为例:7B 参数的 BF16 权重约 14 GB,BF16 梯度约 14 GB,FP32 主权重和两份动量约 84 GB,一个可训练模型合计约 112 GB 逻辑状态,实际会被 ZeRO/FSDP 分片。再训练一套同规模 critic,等于再背一次参数、梯度、优化器状态和激活。删掉 critic 的收益很大,但不能机械地说“所有训练都便宜十倍”:GRPO 每个 prompt 要在线生成一组回答,DeepSeekMath 的实验组大小是 64,采样和奖励打分可能成为新的主成本。 DPO 与 GRPO 给了两条不同的减法路线。DPO 把带 KL 约束的 RLHF 目标改写成偏好二分类,训练时不再在线采样,也不需要显式奖励模型;GRPO 保留在线探索和奖励信号,用同一道题的组内相对分数替代 critic。二者名字相邻,解决的却不是同一个问题。 图 1:DPO 省掉显式奖励模型和在线 RL 循环,GRPO 省掉 critic。Policy 与 reference 是否共享底座、reward 是规则还是模型,会继续改变真实成本。 02. 最小可用理解 先记住五句话: DPO 吃离线偏好对:同一 prompt 下给一个 chosen 和一个 rejected,直接提高 chosen 相对 reference 的概率优势。 GRPO 吃在线样本组:同一 prompt 采样 $G$ 个回答,打分后用组均值当 baseline,不训练 value model。 GRPO 沿用 PPO 的概率比与 clip:它只换了优势估计,没有丢掉“旧策略采样、新策略更新”这套框架。 reference 不是 critic:reference 负责限制长期漂移,critic 负责估计当前状态的未来回报;删掉 critic 后 reference 通常还在。 组内相对分数只回答谁更好:如果一组答案同分,GRPO 在这道题上没有方向;如果奖励尺度跨题不一致,组内归一化反而能消掉尺度差。 因此,从 DPO 到 GRPO 不是把一个损失函数逐项变形为另一个,而是从“离线比较学习”切换到“在线生成、在线评价”。选择算法时先问数据来自哪里,再问显存够不够。 03. 数学推导 3.1 PPO-RLHF 的起点 对 prompt $x$ 和回答 $y$,标准 KL 正则化目标可以写成: $$\max_{\pi}\ \mathbb{E}_{y\sim\pi(\cdot\mid x)}[r(x,y)]-\beta D_{KL}[\pi(\cdot\mid x)\|\pi_{ref}(\cdot\mid x)]$$ $\pi$ 是待训练策略,$\pi_{ref}$ 是固定参考策略,$r(x,y)$ 是回答奖励,$\beta$ 控制策略离参考模型多远。PPO 用采样轨迹、critic 优势和裁剪目标近似优化它。这里至少要维护 actor、critic 和 reference;若奖励由神经网络给出,还要运行 reward model。 3.2 DPO 如何把奖励塞回策略 上面的目标对每个 $x$ 都有闭式最优策略: $$\pi^*(y\mid x)=\frac{1}{Z(x)}\pi_{ref}(y\mid x)\exp\left(\frac{r(x,y)}{\beta}\right)$$ $Z(x)$ 是归一化常数。把式子移项,可把奖励写成策略与参考策略的对数概率比: $$r(x,y)=\beta\log\frac{\pi^*(y\mid x)}{\pi_{ref}(y\mid x)}+\beta\log Z(x)$$ 偏好数据给出同一 prompt 下的胜者 $y_w$ 与败者 $y_l$。用 Bradley-Terry 模型表示“胜者更好”的概率: $$P(y_w\succ y_l\mid x)=\sigma(r(x,y_w)-r(x,y_l))$$ 把上一式代入后,同一道题共享的 $\beta\log Z(x)$ 自动相消。用可训练策略 $\pi_\theta$ 代替未知最优策略,得到 DPO 损失: $$\mathcal{L}_{DPO}=-\mathbb{E}\log\sigma\left(\beta\left[\log\frac{\pi_\theta(y_w\mid x)}{\pi_{ref}(y_w\mid x)}-\log\frac{\pi_\theta(y_l\mid x)}{\pi_{ref}(y_l\mid x)}\right]\right)$$ 括号里有两层差。第一层是 chosen 与 rejected 的差;第二层是当前 policy 相对 reference 的差。只让 policy 提高 chosen 概率还不够,它提高的幅度必须超过 reference 原本已有的偏好。 DPO 的省钱来自训练协议:偏好对提前收集好,训练循环只做似然计算和二分类,不需要当前策略在线生成回答,也不需要在循环里调用显式 reward model。但代价也在这里:数据分布被固定。策略更新后会产生哪些新错误,旧偏好集未必覆盖。 3.3 PPO 的 critic 到底提供了什么 PPO 的优势通常来自 critic: $$A_t=Q(s_t,a_t)-V_\phi(s_t)$$ $V_\phi(s_t)$ 预测从当前 token 前缀继续生成的平均回报。减去它不会改变策略梯度的期望,却能降低方差。问题是,回答只在末尾拿到 $r(x,y)$,critic 却要在每个位置给出估计;它要学习的目标本身就很噪。 如果暂时不追求每个 token 的状态价值,而是对同一道题一次采 $G$ 个完整回答,就能用样本组构造另一种 baseline。 3.4 GRPO 的组内相对优势 从旧策略 $\pi_{old}$ 对同一 prompt $x$ 采样回答 $\{y_1,\ldots,y_G\}$,奖励为 $\{r_1,\ldots,r_G\}$。Outcome supervision 下,GRPO 定义: $$\hat A_i=\frac{r_i-\mathrm{mean}(r_1,\ldots,r_G)}{\mathrm{std}(r_1,\ldots,r_G)+\epsilon}$$ $\epsilon$ 防止标准差为零时除零。均值把“这道题本来容易还是难”扣掉,标准差把不同题目的奖励尺度拉到接近一致。对一条回答里的每个有效 token,原始 GRPO 都广播同一个 $\hat A_i$:得分高于组均值的回答整体增概率,低于均值的整体降概率。 图 2:两道题的原始 reward 尺度相差很大,组内标准化后各自在零均值参照系里比较。被保留的是同题回答间的相对好坏。 组均值包含样本自身,并不满足“baseline 与本次动作独立”的条件。对固定 prompt 的独立同分布样本,若只减组均值而不除标准差,有 $E[(r_i-\bar r)\nabla\log\pi(y_i)]=(1-1/G)\nabla J$;leave-one-out 均值可去掉这项缩放。再除随机组标准差会改变样本及 prompt 权重,因此完整 GRPO 不能宣称是原始期望奖励梯度的无偏估计。 这不是免费午餐。$G$ 太小时,组均值和标准差很不稳定;奖励全相等时,优势全部为零;同一回答的所有 token 共用一个结果优势时,算法知道“整条回答好”,却不知道哪一步推理真正立功。 3.5 GRPO 仍然是 PPO 家族 对回答 $y_i$ 的第 $t$ 个 token,定义新旧策略概率比: $$\rho_{i,t}(\theta)=\frac{\pi_\theta(y_{i,t}\mid x,y_{i,<t})}{\pi_{old}(y_{i,t}\mid x,y_{i,<t})}$$ GRPO 沿用 PPO-Clip 的保守替代目标,并直接在 loss 上加 reference KL: $$\mathcal{J}_{GRPO}=\mathbb{E}\left[\frac{1}{G}\sum_i\frac{1}{|y_i|}\sum_t\left(\min\left(\rho_{i,t}\hat A_i,\mathrm{clip}(\rho_{i,t},1-\varepsilon,1+\varepsilon)\hat A_i\right)-\beta D_{KL}(\pi_\theta\|\pi_{ref})\right)\right]$$ $\pi_{old}$ 是本轮 rollout 的采样快照,用于纠正同一批数据被重复更新后的分布偏移;$\pi_{ref}$ 是行为锚点,用于限制训练全过程的长期漂移。两者在训练刚开始可能数值相同,但职责不同,不能混用。 DeepSeekMath 使用的单样本 KL 正值估计器是: $$D_{KL}=\exp(\Delta)-\Delta-1,\quad \Delta=\log\pi_{ref}-\log\pi_\theta$$ 因为对任意实数 $\Delta$ 都有 $e^\Delta\geq1+\Delta$,这个量非负。在实现里它可以直接由采到 token 的 policy/reference logprob 计算,不必枚举整个词表。它的期望等于 $D_{KL}(\pi_\theta\|\pi_{ref})$ 的前提是动作从当前 $\pi_\theta$ 采样;实际 rollout 来自 $\pi_{old}$,更新后直接平均得到的是旧采样分布下的代理量,需要相应重要性校正才有同一无偏解释。 3.6 DPO 与 GRPO 的统一观察 DPO 和 GRPO 都在做“相对比较”,但参照物不同: 维度 DPO GRPO 数据 固定的 chosen/rejected 对 当前策略在线生成的一组回答 相对量 policy 相对 reference 的偏好间隔 回答奖励相对同组均值 删除的部件 显式 reward model 与在线 RL 循环 critic/value model 仍需保留 policy、reference、偏好数据 policy、old policy logprob、reference、奖励函数 主要风险 离线数据覆盖不足 采样昂贵、组内同分、奖励投机 这张表也解释了为什么“DPO 之后再做 GRPO”可以成立:先用便宜的离线偏好对把策略推到合理区域,再用在线 GRPO 探索当前策略真正会生成的回答。但这是一种训练安排,不是数学上必须的前后继关系。 3.7 手算一组 DPO 和 GRPO 先看 DPO。设第一对样本中,当前 policy 对 chosen 和 rejected 的整句 logprob 分别为 -2.2 和 -2.8,因此 policy 的偏好间隔是 0.6;reference 对二者的 logprob 分别为 -2.0 和 -2.4,偏好间隔是 0.4。当前策略相对 reference 多出来的间隔只有 0.2。当 $\beta=0.1$ 时,送进 sigmoid 的 logit 是 0.02,胜者概率只略高于 0.5,loss 接近随机猜测的 $\log 2$。 这解释了一个常见现象:chosen 的 logprob 即使很低,也不代表 DPO 一定给它很大梯度。DPO 看的是四个 logprob 组成的相对差。若 reference 已经强烈偏爱 chosen,而 policy 只复制了这种偏好,DPO 不会把它误判成新的进步;若 policy 相比 reference 反而缩小了 chosen 的优势,logit 会变负,损失上升。 再看 GRPO。第一组奖励是 [0.2, 0.8, 1.4, 0.6],均值为 0.75,标准差约为 0.433。标准化后得到 [-1.2699, 0.1154, 1.5008, -0.3463]。第二个回答虽然原始分数 0.8 看上去不高,但它略高于同题平均,因此得到很小的正优势;第三个回答比平均高很多,得到最大的正优势。 假设这四条回答更新后的 token ratio 依次是 [1.2840, 1.1052, 0.7788, 0.9048],裁剪区间为 [0.8,1.2]。第一条优势为负,而 ratio 大于 1,说明坏回答变得更可能,目标必须继续惩罚,不能裁掉。第二条优势为正且 ratio 在区间内,正常提供正向梯度。第三条优势为正但 ratio 低于 0.8,说明好回答被错误降权,目标同样继续惩罚。第四条优势为负且 ratio 在区间内,也正常提供降低概率的梯度。 真正被截住的是“已经朝正确方向走过头”的两种情况:正优势回答的 ratio 超过 1.2,或负优势回答的 ratio 低于 0.8。PPO/GRPO 的 clip 从来不是把所有 ratio 强制塞进区间,而是停止奖励过度乐观的改进。只看 np.clip(ratio) 而不把优势符号放进来,几乎一定会误读目标。 04. 代码实现 code/dpo_grpo_minimal.py 只依赖 NumPy,包含三部分:DPO 二分类损失、GRPO 的组内优势与裁剪损失,以及一个四动作 bandit 的在线训练。实际运行输出如下: 两张解释图由 code/make_figures.py 从文中的固定数值直接生成,图和公式使用同一组奖励,避免手工绘图与代码结果不一致。 === DPO === pair_shape=(2,) preference_logits=[0.02 0.03] loss=0.6807 === GRPO group baseline === rewards.shape=(2, 4) advantages= [[-1.2699 0.1154 1.5008 -0.3463] [-1.2648 -0.6324 0.6324 1.2648]] row_mean=[-0. 0.] === GRPO clipped objective === token_logps.shape=(4, 3) ratio_first_token=[1.284 1.1052 0.7788 0.9048] mean_kl=0.018550, loss=0.1626 奖励张量形状是 [prompt, generation] = [2, 4]。每行优势的均值都为零,说明每道题自己形成 baseline。裁剪例子里第一条优势为负而 ratio 为 1.284,第三条优势为正而 ratio 为 0.7788:二者都朝错误方向越界,因此 min 保留未裁剪项以继续惩罚。另两条 ratio 在区间内。这四条回答都没有进入裁剪收益的平台分支,与 3.7 节一致。 bandit 有四个回答动作,真实奖励分别是 [-0.2, 0.3, 1.0, 0.0]。脚本不训练 critic,每轮只采 32 个动作、组内标准化奖励,再用两轮 PPO 式更新复用这组样本: step P(a0) P(a1) P(a2=best) P(a3) 1 0.2288 0.2494 0.2883 0.2335 20 0.0210 0.0346 0.9179 0.0265 40 0.0048 0.0064 0.9833 0.0055 80 0.0026 0.0035 0.9906 0.0033 第 80 步时,最优动作概率从 0.25 升到 0.9906。这个实验展示的是组 baseline 能提供正确梯度,不代表大模型会如此平滑:真实回答空间巨大,奖励噪声、长度与采样温度都会改变结果。 05. 工业级实现对照 本文人工核对了 Hugging Face TRL 的 GRPOTrainer,以 2026-09-13 的主分支提交 cd2c528 为准。生产实现比论文公式多出几层防护: 奖励函数全返回空值的回答会标成 NaN,从组均值中排除,随后把优势置零,避免“不可评分”被误当成零奖励。 scale_rewards 支持按 group、batch 或不缩放;多目标奖励还可以“先加权求和再归一化”或“各目标先归一化再求和”。两者表达的偏好并不一样。 importance sampling 可以按 token 或 sequence 计算。长回答里 token 级 ratio 连乘会很不稳定,序列级方案则改变了权重粒度。 loss 已不止原始 grpo,还包括 bnpo、dr_grpo、dapo 等不同归一化方式。它们主要在回答长度和 batch/token 归一化上修补偏差。 使用 vLLM 生成时,推理引擎与训练模型的 logprob 可能有数值差异,框架提供额外 importance-sampling correction。 PEFT 场景可在同一底座上挂训练 adapter 与 reference adapter,减少复制完整模型的显存。 最值得监控的日志不是单独一个 reward,而是:组内 reward 标准差为零的比例、clip ratio、KL、回答长度、每种 reward 的均值与方差。如果 reward 上升同时长度暴涨、KL 激增或同分组变多,训练可能在钻评分规则的空子。 工程上还要区分三种 mask:prompt token 不应进入 completion loss,padding token 不应参与平均,工具调用或环境交互产生的非策略 token 也不该伪装成策略动作。mask 错一个维度,代码仍可能正常运行,但 loss 的分母已经变了。 5.1 一次工业 GRPO update 的数据流 第一步是挑 prompt。随机抽题看似公平,但大量过易题会产生“全对组”,大量过难题会产生“全错组”,两者都没有组内方差。生产系统通常会维护题目通过率、奖励方差和最近采样时间,让训练 batch 更多覆盖当前策略刚好有分歧的区域。 第二步是冻结采样快照。系统保存 $\pi_{old}$ 的版本标识,并为每个 prompt 生成 $G$ 条 completion。若用独立 vLLM 服务生成,必须记录推理引擎返回的 token id、mask 和 logprob;只保存文本再重新 tokenize,特殊 token、空格或聊天模板差异都可能让动作序列错位。 第三步是评分。规则验证器适合数学答案、代码测试和工具任务;神经 reward model 适合帮助性、风格和主观质量。多个奖励合并时要先决定是“先加权再组内归一化”,还是“每种奖励各自归一化再加权”。前者保留权重指定的绝对尺度,后者让每种奖励先取得相近话语权。 第四步计算 group statistics。分布式训练里,同一 prompt 的 $G$ 条回答可能散在不同 GPU 上,均值和标准差必须在完整组上聚合。若每张卡只看本地子组,baseline 会随数据切分改变;同一实验换一个 GPU 数就可能得到不同梯度。 第五步计算当前 policy 与 reference 的 token logprob。old_logprob 必须保持 rollout 时的值,不能在每个优化 epoch 重算;reference_logprob 可以预计算并缓存,也可以通过冻结模型或禁用 LoRA adapter 得到。三种 logprob 看着形状相同,来源却完全不同。 第六步形成 per-token loss。Outcome reward 把回答级优势广播到有效 completion token,再乘 ratio、做 clip、加 KL。随后按 mask 归一化。先对每条序列取平均再对 batch 平均,会让长短回答权重相近;直接对全 batch token 求平均,则长回答贡献更多。两种都能运行,但优化目标不同。 第七步更新并监控。一次 rollout 可做有限个优化 epoch,随后必须重新采样。应同时记录 reward、组内标准差、zero-std 比例、KL、clip fraction、回答长度和吞吐量。只有 reward 上升而这些量失控时,不能把结果叫作训练成功。 5.2 六个不会报错却会训歪的实现细节 分组边界错位。 假设 batch 排列本应是“题目 A 的八条回答、题目 B 的八条回答”,数据加载器却在聚合前打乱了 completion,view(-1, G) 仍然合法,均值却混合了不同题目。最可靠的做法是保留 prompt id,计算优势前断言每个连续组的 id 完全相同,分布式 gather 后再断言每题正好有 $G$ 条记录。 把样本标准差和总体标准差混着用。 NumPy 的 std 默认分母是 $G$,PyTorch 某些接口曾默认用 $G-1$。组很大时差别不明显,$G=2$ 或 4 时会明显改变优势幅度。公式、最小实现和训练框架必须约定同一个定义,日志里也应记录原始 reward,而不是只保留标准化后的值。 先逐卡归一化再全局聚合。 数据并行下每张卡只拿到同一题的一部分回答,如果先在本地求均值,结果取决于回答被分到哪张卡。正确顺序是先收集完整组的 reward,再统一计算统计量,最后切回各进程需要的那一段优势。 生成文本与训练 token 不一致。 推理服务可能自动补聊天模板、合并空白、提前截断,训练端重新编码字符串后得到另一串 token。此时旧策略 logprob 与当前策略 logprob 对应的动作不同,ratio 已经失去意义。rollout 应把 token ids 视作主数据,文本只用于评分和审计;恢复训练时也要保留 tokenizer 与 chat template 版本。 长度归一化悄悄改变偏好。 若每条回答先按自身长度求 token loss 均值,短回答和长回答各占一个样本权重;若整个 batch 按 token 总数求均值,长回答天然贡献更多。再叠加“回答越长越可能碰到一次正确步骤”的 reward,模型可能学会延长推理。应同时画 reward 对长度的散点图,并用长度分桶比较通过率。 奖励器和策略一起漂移。 在线迭代训练 reward model 时,今天的 0.8 与上周的 0.8 可能不是同一尺度。组内归一化能消掉部分仿射漂移,却消不掉排序规则改变。需要保留固定校准集和冻结的旧奖励器,周期性比较新旧排序一致率;否则策略看似持续提高,其实只是追着变化的裁判跑。 最后还要保存可恢复状态:policy、optimizer、scheduler、reference 版本、rollout 生成参数、reward 版本和数据游标缺一不可。只保存 actor 权重可以用于推理,却无法保证训练恢复后仍在优化同一个目标。GRPO 少了 critic checkpoint,但并没有把实验追踪简化成一个模型文件。 06. 代价与边界 删掉 critic,省下的是真实成本。 全参数训练时少一套同规模可训练状态和激活;LoRA 或 ZeRO 下绝对节省会变化,但工程链路仍少了 value target、value loss、value clipping 与 critic checkpoint。 组采样把成本转移到生成端。 GRPO 必须对同一 prompt 采多个回答。推理较长、组大小较大或 reward model 很重时,rollout 可能比 critic 更贵。适合有便宜可验证奖励的数学、代码和工具任务;对主观写作、开放式对话,可靠打分本身就是难题。 组内归一化丢掉跨题绝对难度。 一道题的最高分只有 0.2,另一道题四个答案都超过 0.9,标准化后仍可能产生相似幅度的优势。前者也许是奖励器不可靠或题目不可解,却会得到同等更新权重。 同分会让梯度消失。 二值正确性奖励在策略已经很强或很弱时尤其明显:一组全对或全错,标准差接近零。这时增加组大小、提高采样温度、设计更细奖励或主动挑选处在学习边界的 prompt,比盲目调学习率更有效。 Outcome reward 的信用分配很粗。 一条 2000-token 推理只拿一个终局优势,每个 token 被同方向推动。过程奖励可以细化信号,但会引入步骤标注、奖励模型偏差与“奖励器教模型写格式”的新问题。 reference KL 不能修复坏奖励。 KL 只限制走多远,不判断方向对不对。奖励函数偏爱冗长、套话或特定格式时,较小 KL 只是让投机变慢。 6.1 什么时候选 DPO,什么时候选 GRPO 如果手里已经有高质量 chosen/rejected 数据,生成和人工标注又很贵,DPO 通常是更稳的第一步。它像普通监督训练一样容易批处理,失败时也容易回查是哪一对偏好造成梯度。它尤其适合风格、安全边界、回答格式这类可以提前整理成对的数据。 如果任务能自动验答案,而且策略必须靠探索发现新解法,GRPO 更合适。数学最终答案、代码单元测试、工具调用结果和可执行环境都能提供相对便宜的在线奖励。此时 critic 很难从稀疏终局分数学习每个 token 的价值,组 baseline 的简化更有吸引力。 如果奖励高度主观、每条回答又很长,两者都不会自动解决问题。DPO 受限于离线偏好覆盖,GRPO 受限于奖励模型可靠性和 rollout 成本。常见组合是 SFT 打基础、DPO 吸收稳定的人类偏好、GRPO 在可验证子任务上继续在线探索,最后再用独立评测检查能力回退。 做预算时至少拆成四项:训练模型状态、冻结模型前向、在线生成 token 数、奖励评估次数。只比较“有几个模型”会漏掉 GRPO 的采样账;只比较“每步生成多少 token”又会漏掉 PPO critic 的反向和优化器状态。把四项分别量出来,才能判断删掉 critic 是否真的让当前系统便宜了一个量级。 还有一笔常被忽略的是通信。PPO 的 critic 若与 actor 分开部署,需要在 rollout、价值推理和更新之间搬运 token、mask、value 与 advantage;若与 actor 共卡,又会争夺显存和计算流。GRPO 把 value 请求换成组统计,统计量只是一组标量,通信链更短。可是一旦奖励模型远程部署,$G$ 条回答仍要跨服务传输并等待评分。优化时应测端到端每步墙钟时间,把生成、评分、训练和等待分别打点;单看 GPU 峰值显存,很容易把“省显存”误写成“省总成本”。 评估也要和训练奖励解耦。至少准备一套不参与更新的题目、一个不同实现的验证器,以及人工抽查样本。训练奖励与独立评测同时上升,才说明策略能力可能真的提高;只有训练奖励上涨,最多能证明优化器成功找到了当前评分函数偏爱的输出。对推理模型还应分别统计最终答案、推理有效性、格式合规和长度,避免一个聚合分数掩盖能力回退。 这些记录也是定位回归、比较算法和复现实验的最低证据。 07. 经典论文脉络 Proximal Policy Optimization Algorithms(arXiv:1707.06347,2017)提出可多轮复用 rollout 的裁剪替代目标,是 GRPO 概率比与 clip 的直接来源。 Direct Preference Optimization(arXiv:2305.18290,2023)利用 KL 正则化 RLHF 的闭式最优策略,把偏好学习变成二分类,训练阶段绕开在线 RL。 DeepSeekMath(arXiv:2402.03300,2024)正式提出 GRPO,用同题多回答的组分数替代 value model;7B 实验中 GSM8K 从 82.9% 提升到 88.2%,MATH 从 46.8% 提升到 51.7%。 DeepSeek-R1(arXiv:2501.12948,2025)把大规模 GRPO 推到长链推理与可验证奖励场景,也让组内相对优化成为推理模型训练的常用基线。 这些数字来自各论文自己的设置,不能跨模型直接比较。DPO 的优势在训练简单,GRPO 的优势在能跟随当前策略在线探索;算法名本身不决定最终能力。 08. 常见误解 “GRPO 不需要奖励模型。” 错。GRPO 不需要 critic,但必须有奖励信号。奖励可以是神经 reward model、规则验证器、单元测试或多个评分器的组合。 “DPO 和 GRPO 都去掉 value,所以是一回事。” DPO 通过重参数化绕开整套在线 RL;GRPO 仍是 policy-gradient 方法,只把优势 baseline 从 learned value 换成 group statistics。 “组均值就是无偏的真实价值函数。” 组均值只是当前策略、当前 prompt、当前采样温度下的 Monte Carlo baseline。它会随组大小和采样分布变化,不是对所有状态都可查询的 $V(s)$。 “归一化后奖励尺度不重要。” 组内仿射变换会被抵消,但排序错误、同分、异常值和多目标权重仍会改变梯度。reward design 依然是核心工作。 “去掉 critic 就一定便宜十倍。” 不成立。删掉一个同规模可训练模型可能显著降低显存和通信,但 GRPO 增加了组采样。总成本取决于序列长度、组大小、reward 推理、并行方式和是否使用 LoRA。 “clip 能保证 KL 很小。” clip 只截断采样动作上的替代收益;共享参数更新仍可能让未采到 token 的分布大幅变化,所以工程实现仍单独监控或惩罚 KL。 09. 动手验证 在仓库目录运行: python3 outputs/fundamentals_files/dpo_grpo/code/dpo_grpo_minimal.py 可以做四个改动观察边界: 把 group_size=32 改成 2,多换几个 seed。预期收敛抖动明显增大,有时一组采不到最优动作。 把 reward_table 改成四个相同值。预期标准差为零,优势接近零,策略不再学习。 把最好动作奖励从 1.0 改成 0.31,使它只比第二名高 0.01。预期需要更多样本才能稳定区分。 把 beta 从 0.02 提到 0.5。预期策略更贴近均匀 reference,最优动作概率上升更慢、上限更低。 再做一个 DPO 对照:交换 policy_chosen 与 policy_rejected,观察 preference logit 变负、loss 变大。它说明 DPO 优化的是相对 reference 的 chosen/rejected 间隔,而不是 chosen 的绝对 logprob。 10. 延伸阅读 读这篇之前建议先看: 策略梯度与 PPO 基础:优势估计、概率比和 clip 的完整推导。 读完这篇可以继续看: 把 RL 用到扩散模型上(DiffusionRL):把多步去噪改写成 MDP 轨迹。 视频生成中的强化学习与奖励模型:GRPO 进入视觉生成后,采样成本、时序奖励与 reward hacking 如何变化。 附录:完整代码 09 节用到的脚本全文如下(dpo_grpo_minimal.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 dpo_grpo_minimal.py """DPO 与 GRPO 的最小可运行实现。 运行:python3 dpo_grpo_minimal.py 只依赖 NumPy,不需要模型权重或外部数据。 """ import numpy as np def log_sigmoid(x): """数值稳定的 log(sigmoid(x))。""" return -np.logaddexp(0.0, -x) def dpo_loss(policy_chosen, policy_rejected, ref_chosen, ref_rejected, beta=0.1): """DPO 二分类损失;四个输入都是整条回答的对数概率。""" policy_gap = policy_chosen - policy_rejected reference_gap = ref_chosen - ref_rejected logits = beta * (policy_gap - reference_gap) return -log_sigmoid(logits).mean(), logits def group_advantages(rewards, eps=1e-4): """对同一 prompt 的 G 个奖励做组内中心化和标准化。""" mean = rewards.mean(axis=1, keepdims=True) std = rewards.std(axis=1, keepdims=True) return (rewards - mean) / (std + eps) def grpo_loss(new_logps, old_logps, ref_logps, advantages, mask, clip_eps=0.2, beta=0.04): """Outcome-supervision GRPO:一条回答的优势广播给它的全部 token。""" log_ratio = new_logps - old_logps ratio = np.exp(log_ratio) clipped_ratio = np.clip(ratio, 1 - clip_eps, 1 + clip_eps) advantage_per_token = advantages[:, None] surrogate = np.minimum( ratio * advantage_per_token, clipped_ratio * advantage_per_token, ) # DeepSeekMath 使用的正值 KL 估计器:exp(x) - x - 1,其中 x=log pi_ref-log pi。 ref_gap = ref_logps - new_logps per_token_kl = np.exp(ref_gap) - ref_gap - 1 per_token_loss = -surrogate + beta * per_token_kl sequence_loss = (per_token_loss * mask).sum(axis=-1) / mask.sum(axis=-1) return sequence_loss.mean(), ratio, per_token_kl def softmax(logits): shifted = logits - logits.max() exp_logits = np.exp(shifted) return exp_logits / exp_logits.sum() def train_group_bandit(steps=80, group_size=32, seed=7): """用 GRPO 训练四选一 bandit,展示无 critic 的在线更新。""" rng = np.random.default_rng(seed) logits = np.zeros(4, dtype=np.float64) reference_logits = np.zeros(4, dtype=np.float64) reward_table = np.array([-0.2, 0.3, 1.0, 0.0]) lr, clip_eps, beta = 0.18, 0.2, 0.02 snapshots = [] for step in range(steps): old_logits = logits.copy() old_probs = softmax(old_logits) actions = rng.choice(4, size=group_size, p=old_probs) rewards = reward_table[actions] advantages = (rewards - rewards.mean()) / (rewards.std() + 1e-4) old_selected = np.log(old_probs[actions]) ref_probs = softmax(reference_logits) ref_selected = np.log(ref_probs[actions]) # 同一组 rollout 做两轮 PPO 式复用;old_logprob 始终来自采样快照。 for _ in range(2): probs = softmax(logits) new_selected = np.log(probs[actions]) ratio = np.exp(new_selected - old_selected) active = ~(((advantages >= 0) & (ratio > 1 + clip_eps)) | ((advantages < 0) & (ratio < 1 - clip_eps))) # d[-min(rA, clip(r)A)]/d log pi,以及采样 KL 项的导数。 coefficient = -active.astype(np.float64) * ratio * advantages ref_gap = ref_selected - new_selected coefficient += beta * (1 - np.exp(ref_gap)) one_hot = np.eye(4)[actions] grad_logp = one_hot - probs[None, :] grad_logits = (coefficient[:, None] * grad_logp).mean(axis=0) logits -= lr * grad_logits if step in {0, 19, 39, 79}: snapshots.append((step + 1, softmax(logits).copy())) return snapshots if __name__ == "__main__": np.set_printoptions(precision=4, suppress=True) policy_chosen = np.array([-2.2, -1.7]) policy_rejected = np.array([-2.8, -2.1]) ref_chosen = np.array([-2.0, -1.9]) ref_rejected = np.array([-2.4, -2.0]) loss_dpo, logits_dpo = dpo_loss( policy_chosen, policy_rejected, ref_chosen, ref_rejected ) print("=== DPO ===") print(f"pair_shape={policy_chosen.shape}") print(f"preference_logits={logits_dpo}") print(f"loss={loss_dpo:.4f}") rewards = np.array([[0.2, 0.8, 1.4, 0.6], [8.0, 9.0, 11.0, 12.0]]) advantages = group_advantages(rewards) print("\n=== GRPO group baseline ===") print(f"rewards.shape={rewards.shape}") print(f"advantages=\n{advantages}") print(f"row_mean={advantages.mean(axis=1)}") old_logps = np.log(np.array([ [0.30, 0.25, 0.20], [0.18, 0.22, 0.20], [0.12, 0.16, 0.20], [0.25, 0.20, 0.18], ])) new_logps = old_logps + np.array([ [0.25, 0.25, 0.25], [0.10, 0.10, 0.10], [-0.25, -0.25, -0.25], [-0.10, -0.10, -0.10], ]) ref_logps = old_logps - 0.05 mask = np.ones_like(old_logps) one_group_adv = group_advantages(np.array([[0.2, 0.8, 1.4, 0.6]]))[0] loss_grpo, ratio, kl = grpo_loss( new_logps, old_logps, ref_logps, one_group_adv, mask ) print("\n=== GRPO clipped objective ===") print(f"token_logps.shape={new_logps.shape}") print(f"ratio_first_token={ratio[:, 0]}") print(f"mean_kl={kl.mean():.6f}, loss={loss_grpo:.4f}") print("\n=== Online group bandit ===") print("step P(a0) P(a1) P(a2=best) P(a3)") for step, probs in train_group_bandit(): print(f"{step:>4} " + " ".join(f"{p:.4f}" for p in probs)) make_figures.py """生成 DPO/GRPO 教程的两张解释图。""" 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, }) def box(ax, x, y, w, h, text, color): patch = FancyBboxPatch( (x, y), w, h, boxstyle="round,pad=0.02,rounding_size=0.03", facecolor=color, edgecolor="white", linewidth=1.5, ) ax.add_patch(patch) ax.text(x + w / 2, y + h / 2, text, ha="center", va="center", fontsize=12, weight="bold", color="#182238") def method_map(): fig, axes = plt.subplots(1, 3, figsize=(14, 4.6), facecolor="#f7f9fc") methods = [ ("PPO-RLHF", [("偏好对", "#dbeafe"), ("奖励模型", "#fde68a"), ("在线采样", "#ddd6fe"), ("Actor + Critic", "#fecaca")], "奖励模型 + critic 都要维护"), ("DPO", [("离线偏好对", "#dbeafe"), ("二分类损失", "#bbf7d0"), ("Policy + Reference", "#ddd6fe")], "省掉显式奖励模型和在线 RL"), ("GRPO", [("同题多回答", "#dbeafe"), ("规则/奖励打分", "#fde68a"), ("组内相对优势", "#bbf7d0"), ("Policy + Reference", "#ddd6fe")], "保留在线采样,省掉 critic"), ] for ax, (title, blocks, note) in zip(axes, methods): ax.set_xlim(0, 1) ax.set_ylim(0, 1) ax.axis("off") ax.set_title(title, fontsize=18, weight="bold", color="#172554", pad=14) gap = 0.08 h = (0.72 - gap * (len(blocks) - 1)) / len(blocks) y = 0.85 - h for text, color in blocks: box(ax, 0.12, y, 0.76, h, text, color) y -= h + gap ax.text(0.5, 0.03, note, ha="center", va="bottom", fontsize=11, color="#475569") fig.suptitle("三种对齐路线省掉的部件不同", fontsize=22, weight="bold", color="#0f172a", y=1.02) fig.tight_layout() fig.savefig(OUT / "alignment_method_map.png", bbox_inches="tight", facecolor=fig.get_facecolor()) plt.close(fig) def group_relative_plot(): rewards = np.array([[0.2, 0.8, 1.4, 0.6], [8.0, 9.0, 11.0, 12.0]]) advantages = (rewards - rewards.mean(1, keepdims=True)) / rewards.std(1, keepdims=True) fig, axes = plt.subplots(1, 2, figsize=(12, 4.8), facecolor="#f7f9fc") colors = ["#60a5fa", "#a78bfa", "#34d399", "#fb7185"] x = np.arange(4) for row, name in enumerate(["题目 A", "题目 B"]): axes[0].bar(x + (row - 0.5) * 0.18, rewards[row], width=0.18, label=name, color=colors[row]) axes[0].set_title("原始奖励不可跨题直接比较", fontsize=16, weight="bold") axes[0].set_xticks(x, [f"回答 {i+1}" for i in x]) axes[0].set_ylabel("reward") axes[0].legend(frameon=False) width = 0.34 axes[1].bar(x - width / 2, advantages[0], width=width, label="题目 A", color="#60a5fa") axes[1].bar(x + width / 2, advantages[1], width=width, label="题目 B", color="#a78bfa") axes[1].axhline(0, color="#334155", linewidth=1) axes[1].set_title("组内标准化只保留相对名次", fontsize=16, weight="bold") axes[1].set_xticks(x, [f"回答 {i+1}" for i in x]) axes[1].set_ylabel("advantage") axes[1].legend(frameon=False) for ax in axes: ax.spines[["top", "right"]].set_visible(False) ax.grid(axis="y", alpha=0.18) fig.suptitle("GRPO:每道题自己组成一个参照系", fontsize=21, weight="bold", color="#0f172a", y=1.02) fig.tight_layout() fig.savefig(OUT / "group_relative_advantage.png", bbox_inches="tight", facecolor=fig.get_facecolor()) plt.close(fig) if __name__ == "__main__": method_map() group_relative_plot() print(OUT / "alignment_method_map.png") print(OUT / "group_relative_advantage.png") 更多 AIGC 论文解读,关注微信公众号「人工智能炼丹君」 每日更新 · 论文精选 · 深度解读 · 技术脉络 微信搜索 人工智能炼丹君 或扫描下方二维码关注
2026年09月13日
4 阅读
0 评论
0 点赞
2026-09-08
AIGC 基本功|策略梯度与 PPO 基础-PPO
策略梯度与 PPO 基础 所属方向:对齐与强化学习 | 难度:核心必修 | 前置知识:无 关键词:策略梯度、REINFORCE、优势函数、GAE、PPO、裁剪、KL 约束 01. 为什么需要它 先看一个很容易把模型训坏的场景。 你有一个已经做过监督微调的语言模型,给同一个问题采样了若干回答,奖励模型认为其中一条更好。最直接的想法是:让模型提高这条回答的概率。于是你把它的对数概率乘上奖励,连续更新十轮。训练日志里的 reward 不断上升,重新采样时却发现模型开始反复输出奖励模型偏爱的句式,内容变长,语言能力下降,最后连原来会答的问题也答不好。 问题不只在奖励。至少有三笔账同时失控了: 信用分配:一个回答只得到一个总分,到底哪些 token 值得鼓励? 估计噪声:这条回答得 8 分,是动作真的好,还是题目本来就容易? 更新幅度:数据由旧策略采出,新策略反复使用同一批数据后已经变了,为什么还能继续把旧结论当真? 策略梯度解决第一笔账:把“提高期望奖励”变成对可采样策略的梯度。价值函数和优势估计解决第二笔账:把状态本身的难易扣掉,只保留“这个动作比预期好多少”。PPO 的裁剪目标处理第三笔账:允许一批昂贵样本做多轮更新,同时阻止单个样本把新旧概率比推得过远。 这三层关系决定了后面的算法谱系。DPO 改写优化问题,GRPO 换掉价值网络,RLOO 换基线,很多 RLHF 系统再加入参考模型 KL、奖励归一化和长度处理。名字不断变,问题仍可追溯到三个量:优势怎么估、概率比怎么算、策略走多远。 02. 最小可用理解 先记住四句话: 策略梯度会提高“优势为正”的动作概率,降低“优势为负”的动作概率。 优势 $A(s,a)$ 是采取动作 $a$ 后的回报,相对状态 $s$ 下平均预期的超额部分。 GAE 用参数 $\lambda$ 在低方差但依赖 critic 的一步 TD,与高方差但依赖较少的完整回报之间插值。 PPO-Clip 比较新旧策略的动作概率比 $r_t$;一旦更新已经朝有利方向超过 $1\pm\epsilon$,该样本不再继续提供奖励。 把一次语言模型生成看成一条轨迹也很直观:状态是 prompt 加已经生成的 token,动作是下一个 token,策略是词表上的概率分布,终局奖励来自奖励模型。critic 预测“从当前前缀继续生成,平均还能拿多少分”,优势则问“刚才选的这个 token,让结局比平均预期好还是差”。 PPO 的 clip 不是把梯度值裁小,也不是保证 KL 一定不超过某个阈值。它裁的是替代目标里的概率比收益。这个区别是理解所有变体的入口。 03. 数学推导 3.1 从期望回报到策略梯度 考虑一个马尔可夫决策过程。$s_t$ 是时刻 $t$ 的状态,$a_t$ 是动作,$r_t$ 是环境在这一步给出的奖励,$\gamma\in[0,1]$ 是折扣因子。策略 $\pi_\theta(a_t\mid s_t)$ 由参数 $\theta$ 控制。一条长度为 $T$ 的轨迹记作 $\tau=(s_0,a_0,r_0,\ldots,s_T)$,折扣回报为: $$R(\tau)=\sum_{t=0}^{T-1}\gamma^t r_t$$ 训练目标是让轨迹的期望回报最大: $$J(\theta)=\mathbb{E}_{\tau\sim\pi_\theta}[R(\tau)]$$ 轨迹概率由初始状态分布 $p(s_0)$、策略和环境转移 $p(s_{t+1}\mid s_t,a_t)$ 连乘得到: $$p_\theta(\tau)=p(s_0)\prod_{t=0}^{T-1}\pi_\theta(a_t\mid s_t)p(s_{t+1}\mid s_t,a_t)$$ 对 $J$ 求梯度,并使用恒等式 $\nabla p=p\nabla\log p$: $$\nabla_\theta J(\theta)=\mathbb{E}_{\tau\sim\pi_\theta}[R(\tau)\nabla_\theta\log p_\theta(\tau)]$$ 环境转移不依赖 $\theta$,因此它在对数梯度中消失,只剩策略项: $$\nabla_\theta J(\theta)=\mathbb{E}\left[R(\tau)\sum_{t=0}^{T-1}\nabla_\theta\log\pi_\theta(a_t\mid s_t)\right]$$ 这就是 REINFORCE 的骨架。还可以利用“未来动作不能影响过去奖励”,把整条轨迹回报换成从 $t$ 开始的 reward-to-go: $$G_t=\sum_{l=0}^{T-t-1}\gamma^l r_{t+l}$$ 于是每个动作只为它之后的结果负责: $$\nabla_\theta J(\theta)=\mathbb{E}\left[\sum_{t=0}^{T-1}\gamma^t G_t\nabla_\theta\log\pi_\theta(a_t\mid s_t)\right]$$ 这里的 $\gamma^t$ 不能漏:本文从固定初始状态、折扣总回报 $J$ 出发,$G_t$ 的折扣却从当前步重新计起。只有 $\gamma=1$,或把该权重吸收到折扣状态访问分布时,才能省略显式 $\gamma^t$。实际 PPO 常把 rollout 的时间步均匀平均,这是常用替代目标,不应与上面的精确有限时域梯度混同。 如果 $G_t$ 总为正,采到的每个动作都会被提高概率,区别只是力度大小。这种估计虽然无偏,方差却很大。我们需要一个不改变期望梯度的参照物。 3.2 基线为什么可以减 从回报减去任何只依赖状态、不依赖本次动作的基线 $b(s_t)$,期望梯度不变。因为对固定状态 $s$: $$\mathbb{E}_{a\sim\pi_\theta}[b(s)\nabla_\theta\log\pi_\theta(a\mid s)]=b(s)\nabla_\theta\sum_a\pi_\theta(a\mid s)=0$$ 最常见的基线是状态价值函数: $$V^\pi(s_t)=\mathbb{E}_\pi[G_t\mid s_t]$$ 动作价值函数把本次动作也作为条件: $$Q^\pi(s_t,a_t)=\mathbb{E}_\pi[G_t\mid s_t,a_t]$$ 两者之差就是优势函数: $$A^\pi(s_t,a_t)=Q^\pi(s_t,a_t)-V^\pi(s_t)$$ $A>0$ 表示这个动作比该状态下的平均动作好,$A<0$ 表示更差。Actor 通过优势更新策略,critic 通过回归回报学习 $V$,所以这类方法叫 actor-critic。 3.3 从 TD 残差到 GAE 真实 $Q$ 和 $V$ 都不知道,只能估计。最短视的估计是一步 TD 残差: $$\delta_t=r_t+\gamma(1-d_t)V_\phi(s_{t+1})-V_\phi(s_t)$$ $V_\phi$ 是参数为 $\phi$ 的 critic,$d_t\in\{0,1\}$ 表示这一步后轨迹是否真正终止。若 $d_t=1$,就不能再 bootstrap 到下一个状态。critic 准确时,$\delta_t$ 是优势的低方差估计;critic 有系统误差时,它也会把误差直接传给 actor。 把未来多个 TD 残差加进来,会得到不同步数的优势估计。GAE 用指数权重把它们合在一起: $$\hat A_t^{\mathrm{GAE}(\gamma,\lambda)}=\sum_{l=0}^{T-t-1}(\gamma\lambda)^l\delta_{t+l}$$ 实际代码不需要为每个 $t$ 再套一层求和。由上式直接得到从后向前的递推: $$\hat A_t=\delta_t+\gamma\lambda(1-d_t)\hat A_{t+1}$$ 两个极端很关键。$\lambda=0$ 时,$\hat A_t=\delta_t$,只看一步,方差小但强依赖 critic。$\lambda=1$ 且轨迹正确终止时,TD 项会望远镜式相消,得到 $G_t-V(s_t)$,对未来 critic 误差不再敏感,但把后续所有随机奖励都带了进来。常见的 $\gamma=0.99,\lambda=0.95$ 是经验起点,不是定律。 训练 critic 时通常使用 value target: $$\hat V_t^{\mathrm{target}}=\hat A_t+V_\phi(s_t)$$ 实现时要区分 terminated 和因时间上限触发的 truncated。真正终止的状态没有未来价值;时间截断通常仍应 bootstrap。把两者都当 done=1,会在每段 rollout 尾部制造系统性低估。 3.4 为什么需要新旧策略概率比 采样成本高,所以 PPO 会冻结采样策略 $\pi_{\theta_{\mathrm{old}}}$,用同一批轨迹对当前策略 $\pi_\theta$ 做多轮小批量更新。数据来自旧策略,目标却要描述新策略,动作重要性比在固定旧策略状态上校正动作分布;它不校正整条轨迹的状态访问分布,所以以下是局部替代目标: $$r_t(\theta)=\frac{\pi_\theta(a_t\mid s_t)}{\pi_{\theta_{\mathrm{old}}}(a_t\mid s_t)}=\exp(\log\pi_\theta-\log\pi_{\theta_{\mathrm{old}}})$$ 未裁剪的替代目标是: $$L^{\mathrm{PG}}(\theta)=\mathbb{E}_t[r_t(\theta)\hat A_t]$$ 刚开始新旧策略相同,$r_t=1$。如果 $\hat A_t>0$,增大该动作概率会使目标变大;如果 $\hat A_t<0$,减小概率会使目标变大。但在同一批数据上更新多轮后,某些比率可能冲到 1.8、0.2,旧数据已不能可靠描述新策略附近的行为。 3.5 PPO-Clip 的 min 到底裁了哪一边 PPO-Clip 的核心目标只有一行: $$L^{\mathrm{CLIP}}(\theta)=\mathbb{E}_t\left[\min\left(r_t\hat A_t,\mathrm{clip}(r_t,1-\epsilon,1+\epsilon)\hat A_t\right)\right]$$ 这里的 min 取较悲观的收益。它必须与优势的符号一起看: 当 $\hat A_t>0$,把好动作的概率比推到 $1+\epsilon$ 以上不再得分;但若概率比跌到 $1-\epsilon$ 以下,损失仍继续变差,算法不会保护错误方向。 当 $\hat A_t<0$,把坏动作的概率比压到 $1-\epsilon$ 以下不再得分;但若坏动作反而变得更可能,惩罚仍继续加重。 所以 clip 是单侧的乐观收益上限。它不会把所有比率强行夹回区间,也不会保证一次优化后每个样本都满足区间约束。一个 batch 内不同样本共享参数,更新某个样本可能把另一个样本推得更远。 PPO 通常最小化总损失,因此策略项前面带负号,再加 critic 的均方误差和熵奖励: $$\mathcal{L}_{\mathrm{total}}=-L^{\mathrm{CLIP}}+c_v\mathbb{E}_t[(V_\phi(s_t)-\hat V_t^{\mathrm{target}})^2]-c_e\mathbb{E}_t[\mathcal{H}(\pi_\theta(\cdot\mid s_t))]$$ $c_v$ 控制价值损失权重,$c_e$ 控制探索强度,$\mathcal H$ 是策略熵。LLM 对齐还常加入相对参考模型 $\pi_{\mathrm{ref}}$ 的 KL 惩罚: $$r_t^{\mathrm{RLHF}}=r_t^{\mathrm{task}}-\beta\left(\log\pi_\theta(a_t\mid s_t)-\log\pi_{\mathrm{ref}}(a_t\mid s_t)\right)$$ 注意这里有两个不同的“旧模型”。$\pi_{\theta_{\mathrm{old}}}$ 是本轮采样快照,用来计算 PPO ratio;$\pi_{\mathrm{ref}}$ 通常是固定的 SFT 参考模型,用来限制长期漂移。把两者混成一个模型,会把更新稳定性和行为保持两个问题混在一起。 3.6 把公式还原成一轮训练 一轮 PPO 可以按下面的顺序执行。顺序很重要,因为每个张量对应的策略版本不同。 第一步,冻结当前 actor 为 $\pi_{\theta_{\mathrm{old}}}$,用它与环境交互,保存每一步的状态、动作、奖励、old_logprob、critic value、终止标记。语言模型场景还要保存 completion mask,明确哪些位置真的是策略采出的动作。 第二步,在 rollout 末端决定是否 bootstrap。若轨迹真正终止,next_value=0;若只是采样窗口用完,则由 critic 估计最后状态价值。然后从后向前计算 $\delta_t$ 和 GAE。此时优势应视作固定训练标签,actor 更新不能反向穿过 GAE 进入旧 value。 第三步,用 $\hat V_t^{\mathrm{target}}=\hat A_t+V_{\mathrm{old}}(s_t)$ 构造 critic 目标。这里加的是采样时保存的旧 value。如果一边更新 critic 一边用新 value 重算 target,目标本身会追着模型移动。 第四步,打乱 rollout,切成 minibatch。当前 actor 对已经采过的动作重新计算 new_logprob,然后用对数概率之差求 ratio。直接先算概率再做除法容易在大词表和低概率动作上出现下溢。 第五步,分别计算未裁剪和裁剪后的 policy surrogate,逐样本取 min,再按有效动作 mask 求平均。同时计算 value loss、entropy bonus,以及可选的 reference KL。 第六步,反向传播并做梯度范数裁剪。一个 rollout 通常会被重复使用若干 epoch;每轮都用同一份 old_logprob,但 new_logprob 随参数更新而变化。若 approximate KL 超阈值,就提前结束剩余 epoch。 第七步,丢弃这批 on-policy 数据,用更新后的 actor 重新采样。此时新策略成为下一轮的 old policy。PPO 能复用的是“一轮之内的有限次数”,并没有把数据永久变成离线训练集。 用张量形状检查这套流程也很有效。经典控制任务常见 reward/value/advantage/logprob 都是 [T, N],其中 $T$ 是 rollout 长度,$N$ 是并行环境数;语言模型常见 [B, L],其中 $B$ 是回答数,$L$ 是序列长度。只要其中一个量偷偷变成 [B] 并发生广播,训练可能不会报错,却会让整条回答的每个 token 共享一个本不该共享的系数。 3.7 手算一条两步轨迹 用一个最小数字例子把 GAE 和 clip 接起来。设两步奖励为 $[0,1]$,critic 预测为 $[0.4,0.7]$,第二步后真正终止,$\gamma=0.9$。最后一步的 TD 残差是 $\delta_1=1-0.7=0.3$;第一步是 $\delta_0=0+0.9\times0.7-0.4=0.23$。当 $\lambda=0.95$ 时: $$\hat A_0=0.23+0.9\times0.95\times0.3=0.4865$$ 对应的 value target 为 $0.4865+0.4=0.8865$,接近完整 Monte Carlo 回报 $0+0.9\times1=0.9$。若 $\lambda=0$,target 只有 $0.23+0.4=0.63$,它完全相信 critic 对下一状态的 0.7 估计。这几个数把“依赖 critic”和“纳入远期真实奖励”的差别直接展开了。 再设旧策略给第一步动作的概率是 0.25,新策略更新后变为 0.30,则 ratio 为 $0.30/0.25=1.2$。若该动作优势为 0.4865,正好到 $\epsilon=0.2$ 的上边界;继续把概率推到 0.35 时,未裁剪收益会按 ratio=1.4 计算,裁剪收益仍只按 1.2 计算,min 选择后者。若优势改为负数,1.4 一侧不会被保护,因为提高坏动作概率应该继续受到惩罚。 这个例子也说明两个超参数并不独立。$\lambda$ 改变优势的大小和噪声,进而改变样本多久撞上 clip;$\epsilon$ 决定同一批优势可以推动概率多远。只调其中一个而不观察 ratio 分布,往往解释不了训练曲线。 04. 代码实现 4.1 逐步算 GAE code/gae_minimal.py 只依赖 NumPy。它先对一条长度为 4 的固定轨迹从后向前递推,再模拟 20000 条随机轨迹,比较不同 $\lambda$ 下 value target 对真实 $V(s_0)$ 的偏差和方差。固定轨迹的真实输出是: rewards.shape=(4,), values.shape=(4,) lambda=0.00 delta=[0.1445 0.143 0.597 0.2 ] A=[0.1445 0.143 0.597 0.2 ] lambda=0.50 delta=[0.1445 0.143 0.597 0.2 ] A=[0.3858 0.4875 0.696 0.2 ] lambda=0.95 delta=[0.1445 0.143 0.597 0.2 ] A=[0.9734 0.8814 0.7851 0.2 ] lambda=1.00 delta=[0.1445 0.143 0.597 0.2 ] A=[1.0652 0.93 0.795 0.2 ] 同一行的 delta 不随 $\lambda$ 变化,因为它只由一步奖励和 critic 决定;$\lambda$ 改变的是未来残差往前传播的强度。最后一步已经终止,所以四组优势都等于 0.2。 随机实验故意给 critic 加入固定误差,结果是: episodes=20000, true_V0=4.2083, critic_V0=4.5083 lambda mean bias std mse 0.00 3.3133 -0.8950 0.7941 1.4315 0.50 3.8265 -0.3818 0.9201 0.9924 0.95 4.1767 -0.0316 1.8754 3.5181 1.00 4.2130 0.0048 2.2029 4.8529 $\lambda$ 从 0 增大到 1,偏差从 -0.8950 降到接近 0,标准差却从 0.7941 增至 2.2029。这个设定里 $\lambda=0.5$ 的均方误差最低;换一个 critic 误差和奖励噪声,最优点也会移动。这正是“偏差—方差折中”的可观测含义。 图 1:同一随机环境与 critic 误差下,$\lambda$ 改变了偏差和方差的配比。曲线由 make_figures.py 根据 20000 条固定随机种子轨迹生成。 4.2 看懂裁剪的符号 code/ppo_clip_minimal.py 先列出单样本裁剪表。下面四行足以解释 min: ratio A ratio*A clip(ratio)*A min clipped? 1.35 1.0 1.350 1.200 1.200 yes 0.65 1.0 0.650 0.800 0.650 no 1.35 -1.0 -1.350 -1.200 -1.350 no 0.65 -1.0 -0.650 -0.800 -0.800 yes 第一行是好动作已经涨太多,因此收益停在 1.2。第二行是好动作被错误地下调,目标仍取更差的 0.65。第三行是坏动作被错误地上调,惩罚保留 -1.35。第四行才是坏动作下降过多后停止继续奖励。clip 只截断“继续沿正确方向冲得更远”的激励。 图 2:蓝色阴影是 $[1-\epsilon,1+\epsilon]$。正优势在右侧形成平台,负优势在左侧形成平台,另一侧继续保留惩罚。 脚本还构造了一个两臂老虎机:旧策略给好、坏动作各 0.5 概率,同一批数据更新 40 轮。不裁剪时,好动作概率一路升到 0.9492,新旧策略 KL 达 0.8231;PPO-Clip 在第三轮越过 1.2 附近后梯度归零,停在 0.6097,KL 为 0.0246: mode epoch p(good) ratio_good grad KL(old||new) unclipped 3 0.6097 1.2193 0.4890 0.0246 unclipped 40 0.9492 1.8984 0.0990 0.8231 PPO-Clip 3 0.6097 1.2193 0.4890 0.0246 PPO-Clip 40 0.6097 1.2193 0.0000 0.0246 为什么停在 1.2193 而不是精确的 1.2?因为脚本用有限学习率做离散更新,第三步从区间内跨到了区间外,跨过以后才失去梯度。真实神经网络也会出现这种越界,所以还要监控 approximate KL,并在过大时提前停止 epoch。 05. 工业级实现对照 参考实现以 2026-09-08 的上游主分支为准。知识树原来记录的 huggingface/trl/trl/trainer/ppo_trainer.py 已经不在 TRL 当前主分支;当前 TRL 的 trainer 目录以 GRPO、DPO、RLOO 等实现 为主。因此本文把仍在维护、能直接核对 PPO 细节的 Stable-Baselines3 PPO.train 作为代码锚点,同时用 TRL GRPOTrainer 观察 LLM 对齐算法如何改写这些组件。 最小脚本与工业实现的差别主要有六处: rollout buffer:工业实现保存 observation、action、old log-prob、value、reward、termination mask,并在采样结束后统一算 GAE。old log-prob 必须冻结,不能在每个 epoch 重算。 优势标准化:常见实现按 minibatch 或整批做 $(A-\mu)/(\sigma+\varepsilon)$。它通常改善数值尺度,却会让一个样本的梯度依赖同批其他样本;batch 太小还可能产生不稳定统计量。 多轮 minibatch:Stable-Baselines3 的默认参数明确包含 n_epochs=10、gamma=0.99、gae_lambda=0.95、clip_range=0.2。这些是基准默认值,不应脱离任务直接复制。 价值函数处理:actor 与 critic 常共享 backbone,并用 vf_coef 合并损失;一些实现还裁剪 value 更新。value clipping 的尺度受奖励缩放影响,不能把 policy clip 的 $\epsilon$ 无脑复用。 额外护栏:除了 ratio clip,还会使用梯度范数裁剪、熵正则、approximate KL、target KL early stop、学习率退火和数值异常检查。Stable-Baselines3 的参数说明也明确指出,clip 本身不足以保证更新一定很小。 LLM 的 token mask:prompt token、padding 和 completion token 必须分开。策略损失通常只落在生成 token 上;序列末端奖励要变成 token 级回报,KL 往往逐 token 计算。错误的 mask 会把 padding 当动作,或者让 prompt 本身参与优化。 在 RLHF 里还要额外记录 reward/score、reward/non_score_reward、policy/approxkl、policy/clipfrac、loss/value、熵、响应长度和终止比例。只看总 reward 上升无法判断是任务能力提高、KL 惩罚变化,还是模型学会钻奖励与长度的空子。 5.1 指标应该怎样联读 clipfrac 必须先确认实现口径:常见日志统计 $|r_t-1|>\epsilon$,但真正进入目标平台还要满足“正优势且 ratio 过大”或“负优势且 ratio 过小”。越界比例不等于梯度被截断比例。它长期接近 0,可能说明学习率过小、epoch 太少、优势太弱,也可能只是策略已经接近局部平稳;它突然接近 1,说明新旧动作概率已大幅偏离,应结合优势符号与 KL 判断;错误方向的越界样本仍有纠正梯度。单独追求某个“漂亮”的 clipfrac 没有意义,应与 KL、entropy 和 reward 一起看。 approximate KL 快速增大而 reward 没有改善,通常先检查学习率、更新 epoch、优势尺度和 mask。如果 KL 很小、entropy 很快下降,则可能是分布变化集中在少数关键 token,平均 KL 把局部坍缩稀释了,需要再看 token 级分位数或最大值。 critic loss 下降也不一定是好消息。若 explained_variance 仍接近 0,critic 可能只学会了回报均值;若训练 loss 很低而新 rollout 上的 value error 很高,critic 在记忆同批噪声。此时 actor 收到的优势基线并不可靠。可以检查优势均值、标准差、与 return 的相关性,并把 value 网络的学习轮数与 actor 分开调节。 语言模型还有一个常见组合信号:任务 reward 上升、KL 惩罚绝对值变大、回答长度持续增加、人工质量不升。它常意味着奖励模型偏好长度或某种格式,策略在用分布漂移换分。解决顺序应是先按长度和题型切片评估奖励,再检查 EOS 奖励与 truncation,最后才调整 $\beta$。只增大 KL 系数会把症状压住,却不能修复有偏的反馈。 5.2 一个最小验收清单 正式放大训练前,可以用小 batch 做四个不依赖最终 reward 的验收。新旧模型完全一致时,所有有效位置的 ratio 应接近 1,approximate KL 应接近 0;把优势全设为 0 后,policy loss 对 actor 的梯度应为 0;交换优势正负号后,选中动作的更新方向应反转;只改变 padding 内容而保持 mask 不变时,损失应不变。这四项能抓住大部分 silent broadcasting、旧概率未冻结和 mask 泄漏问题。 06. 代价与边界 PPO 的优点是实现直接、可对 actor 和 critic 使用普通一阶优化器、同一批 rollout 可以复用多个 epoch。它的代价也很具体。 第一,PPO 仍是 on-policy 算法。ratio 能容忍有限的策略变化,却不能把很久以前的 replay data 变成可靠的当前策略数据。生成模型的 rollout 很贵,这也是许多对齐算法试图绕开在线 PPO 的原因。 第二,critic 会占显存、计算和调参预算。GAE 的质量由奖励噪声、value 误差、终止处理共同决定。critic 拟合得太慢,优势带偏;拟合得太快,又可能记住同一批高噪声回报。 第三,clip 是代理目标的启发式约束。它不等价于 TRPO 的显式 trust region,也不保证真实期望回报单调上升。大 batch、较小学习率、较少 epoch 和 target KL 能降低风险,仍需用新 rollout 验证。 第四,奖励尺度会渗透到整个系统。优势标准化能消除一部分尺度问题,critic loss、value clipping、梯度竞争和终止奖励仍受影响。LLM 的序列长度还会改变 token 级 KL 总量,导致长回答受到更大惩罚或获得更多优化机会。 以下情况不适合直接上 PPO:没有可信在线奖励或无法持续采样;动作空间可枚举且能直接做监督式偏好优化;离线数据远多于在线交互预算;安全约束必须严格满足而不能只靠软惩罚;环境极端稀疏奖励且 critic 没有可学习信号。此时应先解决反馈、数据或约束建模,而不是把希望寄托在 clip 上。 调试时也应先分清“估计坏了”还是“优化走远了”。前者常表现为优势高噪声、critic 解释方差差、不同随机种子更新方向不稳定,应检查奖励、bootstrap 和 GAE;后者常表现为前几个 minibatch 正常,随后 ratio、KL、clipfrac 一起升高,应减少学习率、epoch 或启用 KL early stop。两类问题都会造成 reward 抖动,但修复手段完全不同。 还要保留旧策略下的原始 rollout 指标。只比较更新后的 loss 无法回答“策略是否真的更好”;每轮更新后用新策略重新采样,并按任务、长度和难度切片比较回报,才能排除同一批数据上的代理目标过拟合。至少运行多个随机种子,因为一次 rollout 的偶然高分足以让小样本实验得出相反结论。 07. 经典论文脉络 REINFORCE:Williams 在 1992 年的 Simple Statistical Gradient-Following Algorithms for Connectionist Reinforcement Learning 给出基于采样回报和 log-derivative 的经典策略梯度形式,优点是通用,主要问题是方差高。 TRPO:Trust Region Policy Optimization(2015)用 KL 约束和二阶近似限制策略更新,为“旧数据上走多远”给出更严格的处理,但实现复杂。 GAE:High-Dimensional Continuous Control Using Generalized Advantage Estimation(2015)把 TD($\lambda$) 思想用于优势估计,明确控制 bias 与 variance。 PPO:Proximal Policy Optimization Algorithms(2017)用 clip 或 KL penalty 构造易实现的一阶替代目标,让同一批数据可做多轮 minibatch 更新。 InstructGPT:Training language models to follow instructions with human feedback(2022)展示了 SFT、奖励模型和 PPO 组合成语言模型 RLHF 流水线的代表性实践,也让参考模型 KL 成为对齐工程中的常见组件。 这条演进线并不是“新论文淘汰旧论文”。REINFORCE 给出梯度来源,GAE 改善估计,TRPO/PPO 限制更新,RLHF 再把状态、动作、奖励和约束映射到序列模型。后续算法通常只替换其中一到两块。 08. 常见误解 误解一:优势为正就说明奖励为正 优势是相对量。一个回答奖励为 -2,如果当前 prompt 下平均只能拿 -5,它的优势仍可能为正;一个回答奖励为 9,如果该状态平均是 9.5,它的优势反而为负。策略更新比较的是条件基线,不是绝对分数。 误解二:把优势乘常数不会影响训练 纯策略梯度方向不变,但真实 PPO 还有固定学习率、clip、value loss、熵和 KL 项。优势尺度改变后,各损失的相对权重、越过 clip 边界的速度都会改变。优势标准化是训练定义的一部分,不能只当日志美化。 误解三:ratio 超出区间后一定没有梯度 只有“超出区间且方向有利”的样本会被截断。好动作概率下降得太多、坏动作概率上升得太多时,目标仍保留梯度去纠正。04 节的四行表比背公式更可靠。 误解四:PPO clip 就是 KL trust region clip 在采样动作上限制替代收益,KL 衡量整个动作分布的变化。一个低概率 token 的概率可以相对变化很多而对平均 KL 贡献有限;许多小变化也可能累积出较大 KL。工程实现常同时监控 ratio、clip fraction 和 KL。 误解五:episode 截断等于终止 死亡、成功等真实终止意味着后续价值为零;时间上限只是观察窗口结束,底层状态可能仍有价值。两者都清零 bootstrap 会使 rollout 尾部的 value target 偏低。在固定最大生成长度的 LLM 训练中,这一点对应 EOS 与被长度上限截断的区别。 误解六:old policy 和 reference policy 是同一个概念 old policy 是一轮 PPO 的行为策略快照,几轮 minibatch 后就会更新;reference policy 是长期锚点,通常在 RL 开始时冻结。前者服务于重要性采样,后者服务于行为约束。 09. 动手验证 先在仓库根目录运行: python3 outputs/fundamentals_files/policy_gradient/code/gae_minimal.py python3 outputs/fundamentals_files/policy_gradient/code/ppo_clip_minimal.py 接着做三个小改动,每次只改一个变量并记录输出。 把 gae_minimal.py 的 critic_error 全部设为 0。你会看到 $\lambda=0$ 的偏差大幅下降;奖励噪声不变时,小 $\lambda$ 的低方差优势变得更有吸引力。 把奖励噪声标准差从 0.8 改为 0.1。完整 Monte Carlo target 的方差会明显下降,$\lambda=1$ 的代价随之减小。 把 ppo_clip_minimal.py 的学习率从 0.3 改为 0.03。PPO-Clip 会更接近 ratio=1.2 后才停住,说明 clip 边界不是参数投影,离散优化仍会跨界。 最后给脚本加入 target_kl=0.02 的提前停止条件:每轮更新后计算 KL(old||new),超过阈值就停止。比较它与 ratio clip 的停止轮数。你会看到两种护栏观察的是不同量,也会理解工业实现为什么常把它们同时保留。 10. 延伸阅读 读完这篇可以继续看: 从 DPO 到 GRPO:去掉价值网络:比较偏好优化、组内相对优势与 PPO critic,追踪它们分别改了哪一项。 RLHF 工程实践:继续研究 token mask、参考模型 KL、奖励白化、长度偏置与分布式 rollout。 视频生成中的强化学习与奖励模型:把策略梯度和相对优势迁移到扩散与 flow matching 生成过程,观察奖励黑客如何出现。 回到开头那三个问题,现在可以逐项检查任何新对齐算法:它用什么分配信用,用什么基线或优势降低噪声,用什么机制约束策略变化。只要这三项说清楚,新名词就不会遮住算法真正改变的地方。 附录:完整代码 09 节用到的脚本全文如下(gae_minimal.py、make_figures.py、ppo_clip_minimal.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 gae_minimal.py #!/usr/bin/env python3 """用 NumPy 展示 GAE 递推,以及 lambda 的偏差—方差折中。 依赖:numpy。运行:python3 gae_minimal.py """ import numpy as np def generalized_advantage_estimate(rewards, values, dones, next_value, gamma, lam): """返回 TD 残差、GAE 优势和 value target,输入均为 shape [T]。""" rewards = np.asarray(rewards, dtype=np.float64) values = np.asarray(values, dtype=np.float64) dones = np.asarray(dones, dtype=np.float64) advantages = np.zeros_like(rewards) deltas = np.zeros_like(rewards) gae = 0.0 for t in reversed(range(len(rewards))): value_next = next_value if t == len(rewards) - 1 else values[t + 1] not_done = 1.0 - dones[t] deltas[t] = rewards[t] + gamma * not_done * value_next - values[t] gae = deltas[t] + gamma * lam * not_done * gae advantages[t] = gae return deltas, advantages, advantages + values def fixed_trajectory_demo(): rewards = np.array([0.0, 0.0, 1.0, 0.5]) values = np.array([0.40, 0.55, 0.70, 0.30]) dones = np.array([0, 0, 0, 1]) gamma = 0.99 print("=== 固定轨迹 ===") print(f"rewards.shape={rewards.shape}, values.shape={values.shape}") for lam in (0.0, 0.5, 0.95, 1.0): deltas, advantages, targets = generalized_advantage_estimate( rewards, values, dones, next_value=0.0, gamma=gamma, lam=lam ) print( f"lambda={lam:>4.2f} delta={np.round(deltas, 4)} " f"A={np.round(advantages, 4)} target={np.round(targets, 4)}" ) def bias_variance_demo(seed=7, episodes=20000): """比较 GAE value target 对真实 V(s_0) 的偏差、标准差和均方误差。""" rng = np.random.default_rng(seed) gamma = 0.99 reward_means = np.linspace(0.2, 0.9, 8) horizon = len(reward_means) dones = np.zeros(horizon) dones[-1] = 1.0 true_values = np.zeros(horizon) for t in reversed(range(horizon)): future = 0.0 if t == horizon - 1 else true_values[t + 1] true_values[t] = reward_means[t] + gamma * future # 故意给 critic 加一组固定误差,使短视的 TD target 有偏。 critic_error = np.array([0.30, -0.90, 0.35, -0.25, 0.20, -0.15, 0.10, -0.05]) approx_values = true_values + critic_error rewards_batch = rng.normal(reward_means, 0.8, size=(episodes, horizon)) print("\n=== lambda 的偏差—方差折中(估计 V(s_0))===") print(f"episodes={episodes}, true_V0={true_values[0]:.4f}, critic_V0={approx_values[0]:.4f}") print("lambda mean bias std mse") for lam in (0.0, 0.5, 0.95, 1.0): estimates = np.empty(episodes) for i, rewards in enumerate(rewards_batch): _, advantages, targets = generalized_advantage_estimate( rewards, approx_values, dones, next_value=0.0, gamma=gamma, lam=lam ) estimates[i] = targets[0] bias = estimates.mean() - true_values[0] mse = np.mean((estimates - true_values[0]) ** 2) print(f"{lam:>6.2f} {estimates.mean():>8.4f} {bias:>8.4f} {estimates.std():>8.4f} {mse:>8.4f}") if __name__ == "__main__": fixed_trajectory_demo() bias_variance_demo() make_figures.py #!/usr/bin/env python3 """生成本文的 GAE 偏差—方差图与 PPO 裁剪曲线。 依赖:numpy、matplotlib。运行:python3 make_figures.py 图片写入相邻的 figures/ 目录。 """ from pathlib import Path import matplotlib.pyplot as plt import numpy as np from gae_minimal import generalized_advantage_estimate from ppo_clip_minimal import clipped_surrogate OUT_DIR = Path(__file__).resolve().parent.parent / "figures" def gae_tradeoff_figure(seed=7, episodes=20000): rng = np.random.default_rng(seed) gamma = 0.99 means = np.linspace(0.2, 0.9, 8) dones = np.zeros(len(means)) dones[-1] = 1.0 true_values = np.zeros(len(means)) for t in reversed(range(len(means))): true_values[t] = means[t] + gamma * (0.0 if t == len(means) - 1 else true_values[t + 1]) approx_values = true_values + np.array([0.30, -0.90, 0.35, -0.25, 0.20, -0.15, 0.10, -0.05]) reward_batch = rng.normal(means, 0.8, size=(episodes, len(means))) lambdas = np.array([0.0, 0.25, 0.5, 0.75, 0.95, 1.0]) bias, std, mse = [], [], [] for lam in lambdas: estimates = [] for rewards in reward_batch: _, _, targets = generalized_advantage_estimate( rewards, approx_values, dones, 0.0, gamma, lam ) estimates.append(targets[0]) estimates = np.asarray(estimates) bias.append(estimates.mean() - true_values[0]) std.append(estimates.std()) mse.append(np.mean((estimates - true_values[0]) ** 2)) fig, ax = plt.subplots(figsize=(10, 5.3)) ax.plot(lambdas, np.abs(bias), "o-", linewidth=2.5, label="Absolute bias") ax.plot(lambdas, std, "s-", linewidth=2.5, label="Standard deviation") ax.plot(lambdas, mse, "^-", linewidth=2.5, label="Mean squared error") ax.set(xlabel="GAE lambda", ylabel="Error scale", title="GAE: less bias usually costs more variance") ax.grid(alpha=0.25) ax.legend(frameon=False) fig.tight_layout() path = OUT_DIR / "gae_bias_variance.png" fig.savefig(path, dpi=180, facecolor="white") plt.close(fig) return path def ppo_clip_figure(): ratio = np.linspace(0.4, 1.6, 500) fig, axes = plt.subplots(1, 2, figsize=(11, 4.8), sharex=True) for ax, advantage in zip(axes, (1.0, -1.0)): raw, clipped, objective = clipped_surrogate(ratio, np.full_like(ratio, advantage)) ax.plot(ratio, raw, linestyle="--", linewidth=2, label="Unclipped") ax.plot(ratio, objective, linewidth=3, label="PPO objective") ax.axvspan(0.8, 1.2, color="#2f6df6", alpha=0.08) ax.axvline(1.0, color="black", linewidth=1, alpha=0.5) ax.set_title(f"Advantage = {advantage:+.0f}") ax.set_xlabel("new / old probability ratio") ax.grid(alpha=0.2) axes[0].set_ylabel("Per-sample surrogate") axes[0].legend(frameon=False) fig.suptitle("PPO-Clip is one-sided and depends on the advantage sign", fontsize=14) fig.tight_layout() path = OUT_DIR / "ppo_clip_sign.png" fig.savefig(path, dpi=180, facecolor="white") plt.close(fig) return path if __name__ == "__main__": OUT_DIR.mkdir(parents=True, exist_ok=True) for output in (gae_tradeoff_figure(), ppo_clip_figure()): print(f"saved: {output} ({output.stat().st_size / 1024:.1f} KiB)") ppo_clip_minimal.py #!/usr/bin/env python3 """用 NumPy 拆解 PPO-Clip,并在两臂老虎机上复用同一批数据更新多轮。 依赖:numpy。运行:python3 ppo_clip_minimal.py """ import numpy as np def clipped_surrogate(ratio, advantage, epsilon=0.2): ratio = np.asarray(ratio, dtype=np.float64) advantage = np.asarray(advantage, dtype=np.float64) unclipped = ratio * advantage clipped = np.clip(ratio, 1.0 - epsilon, 1.0 + epsilon) * advantage objective = np.minimum(unclipped, clipped) return unclipped, clipped, objective def clipping_table(): ratios = np.array([1.35, 0.65, 1.35, 0.65, 1.10, 0.90]) advantages = np.array([1.0, 1.0, -1.0, -1.0, 0.5, -0.5]) raw, clipped, objective = clipped_surrogate(ratios, advantages) print("=== 单样本裁剪表 ===") print(" ratio A ratio*A clip(ratio)*A min clipped?") for r, a, u, c, o in zip(ratios, advantages, raw, clipped, objective): flag = "yes" if not np.isclose(u, o) else "no" print(f" {r:>4.2f} {a:>4.1f} {u:>7.3f} {c:>7.3f} {o:>6.3f} {flag}") def sigmoid(x): return 1.0 / (1.0 + np.exp(-x)) def bandit_objective(theta, clipped): """旧策略两动作概率各 0.5;动作 1 优势 +1,动作 0 优势 -1。""" p_good = sigmoid(theta) ratios = np.array([p_good / 0.5, (1.0 - p_good) / 0.5]) advantages = np.array([1.0, -1.0]) if clipped: return clipped_surrogate(ratios, advantages)[2].mean() return np.mean(ratios * advantages) def finite_difference_gradient(theta, clipped, h=1e-5): return (bandit_objective(theta + h, clipped) - bandit_objective(theta - h, clipped)) / (2 * h) def bernoulli_kl(old_p, new_p): return old_p * np.log(old_p / new_p) + (1.0 - old_p) * np.log((1.0 - old_p) / (1.0 - new_p)) def optimize_same_batch(clipped, steps=40, learning_rate=0.3): theta = 0.0 snapshots = [] watch = {0, 1, 2, 5, 10, 20, 39} for step in range(steps): grad = finite_difference_gradient(theta, clipped) theta += learning_rate * grad if step in watch: p_good = sigmoid(theta) ratio_good = p_good / 0.5 snapshots.append((step + 1, p_good, ratio_good, grad, bernoulli_kl(0.5, p_good))) return snapshots def repeated_update_demo(): print("\n=== 在同一批旧策略数据上更新 40 轮 ===") print("mode epoch p(good) ratio_good grad KL(old||new)") for mode, clipped in (("unclipped", False), ("PPO-Clip", True)): for epoch, p_good, ratio, grad, kl in optimize_same_batch(clipped): print(f"{mode:<10} {epoch:>5d} {p_good:>7.4f} {ratio:>7.4f} {grad:>7.4f} {kl:>7.4f}") if __name__ == "__main__": clipping_table() repeated_update_demo() 更多 AIGC 论文解读,关注微信公众号「人工智能炼丹君」 每日更新 · 论文精选 · 深度解读 · 技术脉络 微信搜索 人工智能炼丹君 或扫描下方二维码关注
2026年09月08日
7 阅读
0 评论
0 点赞
2026-09-05
AIGC 基本功|序列并行与 Ring Attention-SP
序列并行与 Ring Attention 所属方向:分布式训练 | 难度:高阶 | 前置知识:张量并行与流水线并行、注意力显存账本、FlashAttention 关键词:序列并行、Ring Attention、长序列、通信与计算重叠、Ulysses 01. 为什么需要它 一个视频或长文被编码成 131072 个 token,hidden size 为 8192,仅一个 BF16 的 $[L,d]$ 张量就占 2 GiB。Transformer 每层不只保存一个这样的张量,还要保存归一化输入、Q/K/V、MLP 中间量、dropout 状态等;标准 attention 若显式保存 $[L,L]$ 分数矩阵,单头 BF16 就是 32 GiB,根本无法进入反向。 数据并行沿 batch 切,每个 rank 仍要处理完整序列;张量并行沿 hidden/head 切,序列长度 $L$ 仍完整复制;流水线并行沿层切,某一层的长序列峰值仍在。长上下文把瓶颈推到第四个维度:必须沿 sequence 切 token。 但“每卡拿一段 token”并不自动成立。MLP 和 LayerNorm 对每个 token 独立,本地就能算;self-attention 中,每个 query 必须看到所有 key/value。若 rank 0 只拿前 1/8 的 K/V,它算出的就成了局部 attention,数学结果已经改变。序列并行的核心问题,是怎样在不复制全部长序列、不近似 attention 的前提下交换必要信息。 目前最常见的两类答案是 Ulysses 与 Ring Attention。Ulysses 先按 sequence 分片,再用 all-to-all 把布局变成按 attention head 分片,让每个 rank 在少量 head 上看到完整 sequence;算完后再 all-to-all 变回来。Ring Attention 则让每个 rank 保留本地 query,K/V block 沿环传递,经过 $P$ 轮后本地 query 恰好见过全序列。 即使尚未阅读FlashAttention 文章,因此先补最小背景:FlashAttention 并不近似 softmax,它把 Q/K/V 分块,并用在线 softmax 的最大值、分母和加权和递推,避免把完整 $L\times L$ 注意力矩阵写回 HBM。Ring Attention 把同一个分块思想扩展到多设备:块不只从显存搬到片上 SRAM,也沿网络搬到下一个 rank。 本文预算脚本给出:$L=131072,d=8192,P=8$ 时,一个 BF16 序列张量从每卡 2 GiB 降到 0.25 GiB,理想降低 8 倍;每层每 rank 绕环接收 K/V 的单向数据约 3.5 GiB。显存不是白省的,它被换成了通信与 $P$ 轮依赖。 02. 最小可用理解 三句话建立框架: 序列并行把 $[B,L,d]$ 沿 $L$ 切成 $P$ 份,使 LayerNorm、dropout、MLP 等序列形激活每卡近似降为 $1/P$。 Ulysses 用 all-to-all 在 sequence shard 与 head shard 间转置;Ring Attention 让 K/V block 绕环流动,本地 Q 用在线 softmax逐块累积精确结果。 两者没有把全注意力的 $O(L^2d)$ 算术复杂度变成线性,只是把计算与内存分散到多卡,并用通信重叠争取接近线性扩展。 如果只记一个区别:Megatron 风格的 sequence parallel 主要切 LayerNorm/dropout 等非 TP 区域的激活;Ulysses/Ring Attention 属于长上下文的 attention 并行,真正处理完整 sequence 上的全局注意力。很多系统把后者称为 context parallel,以免名称混淆。 03. 数学推导 3.1 标准 attention 的账本 令 $Q,K,V\in\mathbb{R}^{L\times d}$,为简洁省略 batch 与 head 维。scaled dot-product attention: $$O=\operatorname{softmax}\left(\frac{QK^\top}{\sqrt d}\right)V$$ 分数矩阵 $S=QK^\top/\sqrt d$ 有 $L^2$ 个元素,算力约为 $O(L^2d)$。若将序列均分到 $P$ 个 rank,每卡持有 $Q_r,K_r,V_r\in\mathbb{R}^{L/P\times d}$。逐 token 的 MLP 可以直接本地算,但 attention 的本地输出需要: $$O_r=\operatorname{softmax}\left(\frac{Q_r[K_0;K_1;\ldots;K_{P-1}]^\top}{\sqrt d}\right)[V_0;V_1;\ldots;V_{P-1}]$$ 也就是说,本地 query 仍必须与所有 rank 的 K/V 相遇。区别只在“把所有 K/V 一次聚齐”还是“逐块流过”。 3.2 在线 softmax:Ring 的数学底座 对一行 query,把 key/value 分成若干块。处理前 $t$ 个块后,保存行最大值 $m_t$、指数和 $\ell_t$ 与未归一化加权和 $u_t$。新块分数为 $s$,块内最大值为 $m_b=\max(s)$,更新全局最大值: $$m_{t+1}=\max(m_t,m_b)$$ 为了把旧累计量换到新的指数基准,定义 $\alpha=\exp(m_t-m_{t+1})$,则: $$\ell_{t+1}=\alpha\ell_t+\sum_j\exp(s_j-m_{t+1})$$ $$u_{t+1}=\alpha u_t+\sum_j\exp(s_j-m_{t+1})v_j$$ 处理完所有块后: $$o=\frac{u_T}{\ell_T}$$ 每次最大值变大时,旧分母与旧分子同时乘 $\alpha$,所以比值语义不变。算法只需保存 $m,\ell,u$ 和当前 score tile,不需保存全矩阵。它与减最大值的稳定 softmax 完全等价,而不是近似。 3.3 Ring Attention 的 $P$ 轮 第 $r$ 个 rank 固定本地 $Q_r$,初始持有 $K_r,V_r$。第 0 轮计算本地块,第 1 轮把 K/V 发给下一个 rank并接收上一个 rank 的块,如此循环。第 $t$ 轮 rank $r$ 处理来源 $(r-t)\bmod P$ 的 K/V;$P$ 轮后每个 $Q_r$ 看过所有块。 本地计算量近似: $$C_{\mathrm{rank}}=O\left(\frac{L}{P}\cdot L\cdot d\right)=O\left(\frac{L^2d}{P}\right)$$ 每 rank 需要接收 $P-1$ 次大小约 $2(L/P)d\,b$ 的 K/V,其中 $b$ 是每元素字节数: $$V_{\mathrm{ring}}=2\frac{P-1}{P}Ldb$$ 当 $L$ 与 $P$ 同比例增长、每卡本地 token 数保持不变时,单轮的 Q/KV block 大小、计算量和传输量近似不变,增长的是轮数和每卡总计算量。是否能覆盖通信取决于单轮 block 是否足够大、有效算力与链路带宽;弱扩展本身不会自动改善单轮的重叠条件。若本地 block 太小、网络慢或 kernel 没有异步双缓冲,通信会裸露出来。 训练反向不能只把前向环倒放那么简单:每个 rank 还要计算对本地 Q、流动 K/V 的梯度,并把属于原拥有者的 dK/dV 沿环累积或归还。checkpoint 与重计算还能减少保存量,但增加计算。 3.4 因果 mask 如何穿过环 自回归 attention 要求 query 位置 $i$ 只能看 key 位置 $j\le i$。分块后,根据全局位置判断: K/V block 完全位于 Q block 未来:整块跳过; 完全位于过去:无需 mask; 与 Q block 重叠:块内施加三角 mask。 跳过未来块可省计算,但不同 rank 的有效工作量会不均。某些实现使用 zigzag 或重新排列 token,使每卡同时持有前后位置,平衡 causal attention 工作量。若只看非因果公式估算吞吐,部署到 causal LLM 时可能偏差很大。 3.5 Ulysses 的 all-to-all 转置 设输入布局是 sequence-sharded: $$[B,L/P,H,D_h]$$ $H$ 是 head 数,$D_h$ 是每头维度。Ulysses all-to-all 把它转成: $$[B,L,H/P,D_h]$$ 每卡序列完整、head 只剩 $H/P$,于是能在本地执行普通 attention;输出再 all-to-all 回 sequence shard。它通常要求 $H$ 能被 $P$ 整除,若 GQA 只有很少 KV heads,约束会更紧。Ring 不要求按 head 数切,但要经历 $P$ 个顺序环步骤。 Ulysses 的 collective 容易利用成熟 all-to-all,但跨节点 all-to-all 对网络争用敏感;Ring 只与邻居点对点通信,更贴合环形拓扑,也更容易细粒度重叠。没有一种方法在所有硬件和 head 配置上恒优。 3.6 Megatron Sequence Parallel 不等于长上下文 attention 传统 TP 的 column/row parallel 边界往往让某些激活在 TP rank 上复制。Megatron sequence parallel 把 LayerNorm、dropout 等逐 token 操作沿 sequence 切分,并用 reduce-scatter 与 all-gather 衔接 TP 区域,使这些激活内存近似降为 $1/P$。它与 TP process group 绑定,主要减少非 TP 区域的重复激活。 但 attention 若仍按 head 做 TP,每个 head 仍可能看到完整 sequence。因此配置项 sequence_parallel=true 并不等于已经支持百万 token。需要进一步的 context parallel、Ulysses 或 Ring Attention 才能沿 attention 的上下文维扩展。 一个序列张量的每卡存储随 P 下降,而完整前向绕环接收的 K/V 字节增加并趋于上限。右图是接收量,不重复计发送量,尚未包含反向通信。 04. 代码实现 ring_attention_sim.py 只用标准库,实现完整 softmax attention 与分块在线递推。Q/K/V 均为 $[4,2]$,每个 K/V block 两行,相当于两轮 ring: Q=K=V shape=(4, 2), kv_block=2, ring_steps=2 full[0]=[1.713718770, 1.717687378] ring[0]=[1.713718770, 1.717687378] max_abs_diff=0.000e+00 本组小样本在打印精度内相同;算法等价性来自第 3.2 节的不变式,浮点实现一般仍存在求和顺序引起的舍入差异。脚本没有真正启动多进程;它把 K/V block 依次喂给同一递推器,模拟每个 rank 收到环上块后的本地数学。生产实现还要异步 send/recv、双缓冲、causal mask 与 backward。 sp_budget.py 对 $L=131072,d=8192,H=64,P=8$ 输出: L=131072, d=8192, heads=64, sp=8, dtype_bytes=2 one [L,d] tensor: replicated=2.00 GiB, sequence-sharded/rank=0.25 GiB ideal memory reduction for sequence-shaped activations=8.0x Ulysses head-divisible=True (64//8=8 heads/rank) Ring Attention rounds=8, local query tokens=16384 Ring K/V traffic per rank per layer (one direction, ideal)=3.50 GiB 2 GiB 只是一个张量,不是整层总显存;0.25 GiB 是理想静态 shard,也没算当前通信块、输出、梯度与工作区。3.5 GiB 对应 $2(P-1)Ldb/P$,是每层、每 rank、单方向的 K/V 接收量。实际网络还包含反向与协议开销。 05. 工业级实现对照 以 2026-09 为准,DeepSpeed 的 deepspeed/sequence/layer.py 中 DistributedAttention 用自定义 all-to-all 在 scatter 轴与 gather 轴间转换 Q/K/V,调用 local_attention 后再变回去。源码明确检查总 head 数必须能被 sequence parallel size 整除,并提供 stream 与 overlap handle 管理通信重叠。 这个工业实现比最小公式多处理 batch 维位置、不同 Q/K/V 布局、rotary position embedding、异步 stream、autograd 的反向 all-to-all 与进程组。布局索引错一位可能 shape 仍合法但语义错误,所以系统必须统一采用明确的张量维约定。 Megatron-LM 的模型并行配置中 sequence_parallel 用于并行 LayerNorm 与 dropout;context_parallel_size 则面向上下文切分。当前 Megatron Core 还支持多种 context parallel 通信方式。读配置时不要只看“SP”缩写,要追到具体张量在哪个轴分片以及 attention 内部采取哪种交换。 Ring Attention 的工业 kernel 通常建立两个 K/V buffer:当前块参与计算时,下一个块在独立通信 stream 接收;计算结束交换 buffer。要真正隐藏通信,必须满足: $$T_{\mathrm{attention\ block}}\ge T_{\mathrm{send/recv}}$$ 此外还要确保计算与 NCCL 没争抢同一资源到互相拖慢。仅在代码里使用 async op 不代表 timeline 上已经重叠,必须用 profiler 验证。 FlashAttention 提供单卡块级 IO 优化,Ring/Ulysses 提供多卡数据布局。两者不是竞争关系:每个 rank 的 local attention 往往正由 FlashAttention kernel 完成。前者减少 HBM 读写,后者减少单卡需要常驻的全局序列。 checkpoint 也必须保存并行元数据。若训练时改变 SP/CP 度,参数本身或许未沿 sequence 切,但 optimizer、随机数状态和位置相关缓存可能变化。恢复前应由框架的 distributed checkpoint 层重新映射,而不是直接让每个 rank 读取旧编号文件。 06. 代价与边界 序列并行省的是序列形激活,不一定省参数和优化器状态;它通常要与 ZeRO/FSDP、TP、PP 组合。完整世界规模可能写成 $W=d_p t_p p_p c_p$,其中 $c_p$ 是 context parallel 度。每多一根轴,都新增 process group、拓扑映射与整除约束。 Ring Attention 保留全注意力的二次计算量。卡数增加能把每卡计算降到 $1/P$,但集群总 FLOPs 仍为 $O(L^2d)$;当上下文翻倍、卡数不变时,总计算约四倍。它突破的是单卡内存和 wall-clock 可扩展性,不是算法复杂度。 通信也不是“免费”。Ring 的每 rank 字节量随 $L$ 线性增长且有 $P$ 轮延迟链;Ulysses 有前后 all-to-all,容易受跨节点网络与并发任务干扰。若 $L$ 尚短,通信与布局转换可能比 attention 本身更贵,普通 TP 更合适。 常见边界包括: Ulysses 的 head 或 KV head 不能被 SP 度整除,GQA/MQA 尤其容易受限; causal mask 造成各 block 计算不均,需要 load-balanced ring; position encoding 必须使用全局位置,不能把每 rank 的局部 token 从 0 重新编号; dropout 的随机数要在切分前后保持统计与重算一致,否则并行度改变会造成难解释的差异; variable-length packed sequence 需要携带边界,不能让不同样本互相 attention; 视频 token 在时间、空间上的排列会影响 causal 或局部结构,切块策略不能只按连续内存方便决定。 在低带宽以太网集群上,增加 context parallel 度可能只是把 OOM 变成通信瓶颈。应先用 activation checkpoint、FlashAttention、减小 microbatch 等单机方法,确认仍由单卡序列容量限制,再引入跨设备 SP。 6.1 从单张量扩展到整层激活账本 一个 $[L,d]$ BF16 张量占 $2Ld$ 字节,但训练层内可能同时存在 residual、norm 输入、Q/K/V、attention output 与 MLP 中间值。设所有与序列线性相关且必须保存到反向的张量合计为 $cLd$ 个元素,序列分片后的理想每卡保存: $$M_{\mathrm{linear/rank}}\approx\frac{cLdb}{P}$$ $c$ 由架构、融合 kernel 与 checkpoint 策略决定,不能用固定常数套所有模型。SwiGLU 会产生门控与 value 两个支路;视频 DiT 还可能有条件分支、cross-attention 与调制参数。最可靠方法是从 profiler 的 allocation timeline 识别跨 backward 存活的张量。 attention tile 与通信 buffer 是额外峰值。Ring 双缓冲至少需要当前和下一份 K/V block;FlashAttention kernel 还需要 score tile、softmax statistic 和 workspace。于是: $$M_{\mathrm{peak/rank}}=M_{\mathrm{linear/rank}}+M_{\mathrm{local\ attention}}+M_{\mathrm{double\ buffer}}+M_{\mathrm{runtime}}$$ 当 $P$ 继续增加时,第一项下降,但 buffer、runtime 和最小 kernel workspace 不会等比下降,显存收益逐渐偏离 $1/P$。用总 allocated 除卡数预测极限会过于乐观。 6.2 反向为何比前向更难 前向中,本地 Q 只需依次看过所有 K/V block。反向要得到 $dQ_r$、$dK_j$ 与 $dV_j$。当来源 $j$ 的 K/V 到达 rank $r$ 时,本地可计算它对 $dQ_r$ 的贡献,也会产生属于来源 rank $j$ 的 $dK_j,dV_j$ 部分。后两者必须跨所有 query shard 求和并归还拥有者。 因此每个块需要身份信息,环中既可能传 K/V,也可能传对应梯度累计。为了减少前向保存,backward 常重算局部 score 与概率,依赖前向保存的行最大值和归一化分母。若随机 dropout 参与 attention,重算还必须恢复相同随机 mask。 理论通信量必须同时列前向与反向。只用 $2(P-1)Ldb/P$ 报告 K/V 前向流量会低估训练网络负担;真实实现还包括 dK/dV、可能的 dQ 归约、控制同步和其他模块的 collective。性能分析应以 profiler 中每层总通信字节与裸露时间为准。 6.3 通信重叠需要满足哪些条件 把一个本地 query block 与一个 K/V block 的 attention 计算时间记为 $T_c$,邻居传输一个 K/V block 的时间记为 $T_n=\alpha+S/\beta$,其中 $\alpha$ 是链路延迟,$\beta$ 是有效带宽,$S$ 是消息字节。双缓冲后的理想每轮时间: $$T_{\mathrm{round}}\approx\max(T_c,T_n)$$ 只有 $T_c\ge T_n$,网络才大部分隐藏。减小 block 能降低临时显存,却增加轮数和延迟占比;增大 block 提高 GEMM 效率,却增大 buffer。block size 是计算、带宽、延迟与显存的共同旋钮。 实现还需避免隐式同步。例如在循环中把 GPU 标量取回 CPU、过早 wait 通信 handle、复用尚未完成接收的 buffer,都会把异步流水串行化。正确做法是独立 stream、event 依赖与 ping-pong buffer,并在 GPU timeline 上确认 send/recv 与 attention kernel 真正重叠。 环的物理顺序同样重要。逻辑 rank 相邻却跨交换机,会让每轮走慢链路;合理 ring 应优先沿 NVLink 域,再通过少量跨节点边连接。多维集群还可以分层:节点内用 Ulysses all-to-all,节点间用 Ring,或反过来根据硬件选择,这正是统一序列并行方法的动机。 6.4 长度、head 与并行轴的联合约束 设 context parallel 度为 $c_p$、tensor parallel 度为 $t_p$。若 TP 已把 head 切为 $H/t_p$,Ulysses 再沿 head 分 $c_p$,就需要本地可见 head 数继续可分: $$H\ \bmod\ (t_pc_p)=0$$ 具体约束取决于框架是在同一 head 轴叠加还是采用不同布局。GQA 中应将 $H_{KV}$ 单独代入;只有 8 个 KV head 的模型很难做大 head-parallel。某些系统复制 K/V 以放松整除,但通信和显存模型随之改变。 sequence length 通常也要求能被 $c_p$ 与 block size 整除。不整除时可以 padding 或不等长 shard。padding 实现简单但会浪费二次 attention 计算;不等长 shard 又使通信消息与计算不均。对长度差异大的数据,先按 token 数分桶往往比依赖动态不等长 ring 更稳定。 TP 与 CP 同时使用还会出现多种 collective。attention 投影可能在 TP 组 all-reduce,context attention 在 CP 组 send/recv,层外梯度还在 DP 组 reduce-scatter。若各组在不同 rank 上以不一致顺序发起 collective,可能死锁。框架必须给跨组通信建立确定顺序或独立 stream 依赖。 6.5 视频与多模态序列的特殊问题 视频 token 常按时间、空间高、空间宽展平。连续切 sequence 可能让每个 rank 持有若干完整帧,也可能持有同一帧的空间条带;两种布局对位置编码、局部 attention、数据增强和负载平衡影响不同。若模型有 temporal attention 与 spatial attention 分解结构,最优切分轴也可能随层变化。 全局位置必须从原始坐标生成。RoPE 若把每个 rank 的局部索引重新从 0 开始,不会报 shape 错,却让不同块位置混叠。二维或三维 RoPE 更要携带时间与空间坐标,而不是只传一个线性 offset。跨模态序列还要保持文本、图像、视频 segment 的 mask 与边界。 视频长度和分辨率经常动态变化。同一 global batch 内若 token 数差异大,padding 使最长样本决定所有 rank 工作量;packed sequence 可减少浪费,却要求 block 不跨样本做 attention。调度系统最好按总 token 数而非样本数组 batch,并把实际有效 token 纳入 loss normalization。 因果结构也不总是标准下三角。视频生成可能使用双向视觉 attention、文本到视频 cross-attention 或分块因果时间 mask。Ring 优化中“未来块可跳过”的判断必须来自真实 mask 结构,不能硬编码 LLM 假设。 6.6 正确性验收与故障定位 第一层验收是 shape:每个边界记录 global shape、local shape、分片轴、global offset、padding 与拥有者。第二层是数学:小尺寸 FP64 单卡作为 oracle,逐项比较 forward 和 dQ/dK/dV。第三层是随机性:开启 dropout 与 checkpoint 后,在固定 seed 下比较统计和可复现范围。 输出不一致时,可按轮 dump 每个 rank 处理的 K/V 来源编号。若少一个或重复一个,是 ring 调度;若所有来源齐全但数值错,检查在线 softmax 的旧累计重标定;若仅 causal 错,检查全局位置与块分类;若 forward 对而 backward 错,检查 dK/dV 归还与跨 query shard 求和。 OOM 时不要只减 block。先区分 persistent activation、通信双 buffer、kernel workspace 与碎片。若 activation 占主导,提高 CP 或 checkpoint 有效;若 buffer 占主导,减 block 才有效;若碎片主导,可调分配器或稳定 shape。错误手段可能降低吞吐却完全不改变峰值来源。 性能验收同时做 strong scaling 与 weak scaling。strong scaling 固定全局 $L$,增加卡数,观察端到端时间能否下降;容量弱扩展固定每卡 local token,令 $L\propto P$:精确全注意力的每卡计算 $L^2/P$ 此时随 $P$ 线性增长,不能期望每卡时间不变;应比较实测时间与这一增长基线,以及通信隐藏率。Ring Attention 可随设备数扩展可处理上下文,但总计算也随之增加。 最后把精度也纳入验收。在线 softmax 很稳定,但不同 block 顺序改变浮点加法;BF16、FP8 attention 更明显。应比较最终任务指标、loss 轨迹和梯度,而不是要求多卡输出 bitwise 等同。若误差远超单纯归约顺序,再排查 scale、mask 与布局。 6.7 Ulysses、Ring 与混合方案如何选 先看 head 约束。若本地 Q head 与 KV head 都能被目标 SP 度整除,且集群 all-to-all 带宽高,Ulysses 路径直接,local attention 可复用成熟 kernel。若 KV head 很少、SP 度很大,Ring 不沿 head 切,通常更容易扩展。再看拓扑:全连接高速交换更适合 all-to-all;邻接带宽强、跨组带宽弱时,环形点对点更自然。 再看本地块计算。长序列、大 head dimension 使每轮计算足以隐藏 K/V 传输,Ring 受益;较短序列或大量小 head 时,$P$ 轮延迟可能突出。Ulysses collective 次数少,但消息全局交换可能造成拥塞。决策依据应是每轮计算时间、有效链路带宽和 collective 实测,不是理论总字节一项。 混合方案把设备组织成二维网格。例如节点内 8 卡做 Ulysses,节点间若干组做 Ring;这样 head 整除只限制节点内度数,跨节点使用邻接通信。代价是布局更复杂、process group 更多,Q/K/V 需要在两个维度间保持一致顺序。只有单一方案确实受限时才值得引入。 还要把 TP 纳入考虑。如果 TP 已经切 head,Ulysses 可用 head 更少;有时降低 TP、提高 CP 更适合长序列,有时模型单层容量又要求 TP 不能降。最优并行度会随训练长度变化,短序列预训练与长序列继续训练未必使用同一布局。 6.8 推理场景与训练不同 prefill 处理整段 prompt,计算形态接近训练前向,序列并行能分摊长 prompt attention。decode 每步只有少量新 query,却要读取不断增长的 KV cache;此时瓶颈常是 KV 带宽与跨卡延迟,而非大块 GEMM。训练中优秀的 Ring block,在逐 token decode 里可能太细、同步轮数太多。 推理还涉及请求级并发。可以把不同请求分给数据并行副本,也可把一个超长请求沿 context 分片。前者吞吐高、单请求受单卡长度限制;后者支持超长上下文、占用多个设备。调度器应根据请求长度动态选择,而不是所有请求固定占一个 CP 组。 KV cache 的拥有关系必须稳定。若每 rank 保存一段 sequence,新增 token 的 K/V 放在哪一段、何时重平衡、beam search 如何复制,都要定义。Paged KV cache 与 context parallel 叠加时,逻辑页、物理 GPU 与全局位置形成三层映射,错误可能表现为偶发内容质量下降而非崩溃。 因此本文通信公式主要用于训练或 prefill,不能直接拿来预测 decode tokens/s。推理基准要分别报告 time-to-first-token、inter-token latency、并发吞吐、KV cache 容量和长短请求混合表现。 6.9 上线前检查清单 上线前确认:sequence、head、KV head 与 block 的整除或 padding 规则明确;每个 shard 保存全局位置;causal/packed mask 按全局索引生成;每轮 K/V 来源不重不漏;在线 softmax 保存并正确重标定最大值、分母与分子;backward 的 dK/dV 回到原拥有者;dropout 重算保持随机一致。 系统侧确认 ring 物理顺序符合拓扑,通信与计算 timeline 真重叠,buffer 生命周期无覆盖,所有 rank collective 顺序一致,checkpoint 记录分片元数据,变更 CP 度可以恢复。性能侧同时做 strong/weak scaling,并在真实长度分布而非单一最大长度上测试。 最后保留小尺寸 oracle。每次升级 attention kernel、通信库、RoPE 或 mask 实现,都自动生成随机 Q/K/V,与 FP64 单卡结果比较 forward 和全部梯度;再跑包含因果、变长、GQA、极端分数和非整除长度的边界集。长序列错误代价很高,小 oracle 是最便宜的保险。 6.10 最后的边界:更长不等于更有用 系统能够处理百万 token,只说明容量与运行时间可接受,不保证模型会利用远距离信息。长上下文实验还要检查位置外推、训练长度分布、检索准确率与有效注意范围。并行系统解决“算得出来”,数据与模型设计决定“学得会不会”。 因此扩展长度时同时保留短上下文质量基线,按距离分桶评估信息召回,并报告有效 token 而非 padding 后长度。否则可能用大量 GPU 计算模型并未使用的上下文。 07. 经典论文脉络 FlashAttention(arXiv:2205.14135):用 IO-aware tiling 与在线 softmax 实现精确 attention,避免物化完整分数矩阵,是多卡分块 attention 的本地 kernel 基础。 Reducing Activation Recomputation in Large Transformer Models(arXiv:2205.05198):提出 Megatron sequence parallel 与选择性重计算,减少 TP 区域之外的重复激活。 DeepSpeed Ulysses(arXiv:2309.14509):用 all-to-all 在 sequence 与 head 布局间转换,使超长序列 attention 可扩展。 Ring Attention with Blockwise Transformers(arXiv:2310.01889):让 K/V 块沿环流动并与分块计算重叠,把可处理上下文随设备数扩展。 USP(arXiv:2405.07719):统一 Ulysses 与 Ring 两类序列并行,针对模型结构与网络拓扑组合两者。 这条演进线从“单卡不保存 $L^2$ 矩阵”,走到“多卡不复制 $L$ 激活”,再走到按 head 结构和网络拓扑组合通信方式。算法、kernel 与系统三层缺一不可。 08. 常见误解 误解一:序列并行把 attention 复杂度从二次降成线性。 精确全注意力的集群总计算仍为 $O(L^2d)$;它只把工作分摊,并让内存近似随本地 token 数增长。 误解二:每卡切一段 token 后直接做本地 attention。 那只允许 token 看同一分片,会改变模型。必须让 query 见到全局 K/V,或明确接受局部/稀疏近似。 误解三:Ring Attention 是近似注意力。 在线 softmax 递推在精确算术下与完整 softmax 等价;差别来自浮点归约顺序,不是丢弃连接。 误解四:开启 Megatron sequence_parallel 就解决长上下文。 它主要减少 LayerNorm、dropout 等激活复制;长上下文 attention 通常还需 context parallel。 误解五:异步 send/recv 就等于通信被隐藏。 只有计算时间覆盖传输、stream 真能并发且无资源争抢,timeline 上才会重叠。 误解六:Ulysses 的 SP 度只受 GPU 数限制。 head 数、KV head 数与 tensor layout 必须可切;GQA 模型常比标准 MHA 更早碰到上限。 09. 动手验证 先运行 ring_attention_sim.py,把 block_size 改成 1、2、4。预期 max_abs_diff 始终接近机器精度;再把 score 全部加 1000,稳定递推仍有限,而直接 exp(score) 会溢出。这同时验证精确性与减最大值的重要性。 为脚本加入 causal mask,使用 query/key 的全局索引屏蔽 $j>i$。预期分块结果与完整下三角 attention 一致。然后故意用局部索引 mask,会发现后续 rank 错误允许关注未来或错误屏蔽过去,这是真实分布式实现中很隐蔽的 bug。 运行 sp_budget.py,固定每卡 token 为 16384,让 $L$ 与 ranks 同比例从 1、2、4、8 增长。每卡序列张量保持 0.25 GiB,而 ring K/V 流量与全局 $L$ 增长。这个弱扩展实验说明“可放下更长序列”并不等于“通信不增长”,只是计算增长可能覆盖通信。 在真实集群比较 FlashAttention 单卡、Ulysses、Ring/context parallel:固定模型与 global sequence,记录峰值显存、attention kernel 时间、collective/点对点时间、重叠率与 MFU。分别在节点内和跨节点运行;预期最佳方法随 NVLink、InfiniBand、head 数和 local block 大小变化。 最后做数值验收:关闭 dropout,用同一 Q/K/V 比较单卡与多卡输出及 dQ/dK/dV。BF16 下用合理容差,同时单独检查每个 rank 拼接后的全局顺序。输出对而梯度错,通常说明 backward 中流动 K/V 梯度没有正确归还拥有者。 10. 延伸阅读 读懂本文的完整前置路径是: 数据并行与 ZeRO 显存切分:先区分状态分片与计算分片(已发布)。 张量并行与流水线并行:理解 process group、collective、microbatch 与拓扑(已发布)。 自注意力机制的计算与显存账本:理解 $L^2$ 分数矩阵从哪里来。 FlashAttention 为什么不需要存下注意力矩阵:深入单卡在线 softmax 与 IO 复杂度。 DDPM 训练目标与采样流程和 DiT:用 Transformer 替掉 UNet:理解视频 DiT 的时空 token 为何迅速推高注意力成本。 选 Ulysses 还是 Ring,不应从名字出发。先列出 $L,H,H_{KV},d$、网络拓扑、每卡内存与 causal 结构;再计算 local shape、通信字节、轮数和可重叠窗口;最后用 profiler 验证。长上下文系统真正的能力,是让数学等价、内存边界和物理网络三者同时对齐。 附录:完整代码 09 节用到的脚本全文如下(ring_attention_sim.py、sp_budget.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 ring_attention_sim.py #!/usr/bin/env python3 """Exact blockwise attention recurrence used by ring-style attention; stdlib only.""" import math def dot(a, b): return sum(x * y for x, y in zip(a, b)) def full_attention(q, k, v): outputs = [] scale = math.sqrt(len(q[0])) for query in q: scores = [dot(query, key) / scale for key in k] peak = max(scores) weights = [math.exp(score - peak) for score in scores] denom = sum(weights) outputs.append([sum(w * value[j] for w, value in zip(weights, v)) / denom for j in range(len(v[0]))]) return outputs def blockwise_attention(q, k, v, block_size): scale = math.sqrt(len(q[0])) peak = [-math.inf] * len(q) denom = [0.0] * len(q) numer = [[0.0] * len(v[0]) for _ in q] for start in range(0, len(k), block_size): kb, vb = k[start:start + block_size], v[start:start + block_size] for row, query in enumerate(q): scores = [dot(query, key) / scale for key in kb] new_peak = max(peak[row], max(scores)) correction = math.exp(peak[row] - new_peak) if peak[row] != -math.inf else 0.0 weights = [math.exp(score - new_peak) for score in scores] denom[row] = correction * denom[row] + sum(weights) numer[row] = [correction * old + sum(w * value[j] for w, value in zip(weights, vb)) for j, old in enumerate(numer[row])] peak[row] = new_peak return [[x / denom[row] for x in numer[row]] for row in range(len(q))] q = [[1.0, 0.0], [0.0, 1.0], [1.0, 1.0], [-1.0, 1.0]] k = [[1.0, 0.0], [0.0, 1.0], [1.0, 1.0], [1.0, -1.0]] v = [[1.0, 2.0], [2.0, 0.0], [0.0, 3.0], [4.0, 1.0]] full = full_attention(q, k, v) streamed = blockwise_attention(q, k, v, block_size=2) max_diff = max(abs(a - b) for ra, rb in zip(full, streamed) for a, b in zip(ra, rb)) print("Q=K=V shape=(4, 2), kv_block=2, ring_steps=2") print("full[0]=[" + ", ".join(f"{x:.9f}" for x in full[0]) + "]") print("ring[0]=[" + ", ".join(f"{x:.9f}" for x in streamed[0]) + "]") print(f"max_abs_diff={max_diff:.3e}") sp_budget.py #!/usr/bin/env python3 """Compare idealized per-rank sequence-memory and method constraints.""" import argparse parser = argparse.ArgumentParser() parser.add_argument("--tokens", type=int, default=131072) parser.add_argument("--hidden", type=int, default=8192) parser.add_argument("--heads", type=int, default=64) parser.add_argument("--ranks", type=int, default=8) args = parser.parse_args() if args.tokens % args.ranks: parser.error("tokens must be divisible by ranks") tensor_gib = args.tokens * args.hidden * 2 / 1024**3 local_gib = tensor_gib / args.ranks print(f"L={args.tokens}, d={args.hidden}, heads={args.heads}, sp={args.ranks}, dtype_bytes=2") print(f"one [L,d] tensor: replicated={tensor_gib:.2f} GiB, sequence-sharded/rank={local_gib:.2f} GiB") print(f"ideal memory reduction for sequence-shaped activations={args.ranks:.1f}x") print(f"Ulysses head-divisible={args.heads % args.ranks == 0} ({args.heads}//{args.ranks}={args.heads // args.ranks} heads/rank)") print(f"Ring Attention rounds={args.ranks}, local query tokens={args.tokens // args.ranks}") print(f"Ring K/V traffic per rank per layer (one direction, ideal)={2*tensor_gib*(args.ranks-1)/args.ranks:.2f} GiB") make_figures.py """Regenerate this article's deterministic teaching figure (numpy + matplotlib).""" from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np OUT=Path(__file__).resolve().parents[1]/"figures" OUT.mkdir(exist_ok=True) plt.rcParams.update({"font.sans-serif":["PingFang SC","Arial Unicode MS","DejaVu Sans"],"axes.unicode_minus":False}) p=np.array([1,2,4,8,16,32]); L=131072;d=8192;b=2;g=1024**3 fig,axes=plt.subplots(1,2,figsize=(10,4.4)) axes[0].plot(p,np.full(p.shape,L*d*b/g),"--",label="Replicated") axes[0].plot(p,L*d*b/p/g,"o-",label="Sequence shard") axes[0].set(xscale="log",xlabel="Ranks P",ylabel="One tensor per rank (GiB)");axes[0].legend() axes[1].plot(p,2*(p-1)/p*L*d*b/g,"o-") axes[1].set(xscale="log",xlabel="Ranks P",ylabel="Forward K/V received per rank (GiB)") for ax in axes:ax.grid(alpha=.25) fig.suptitle("Ring attention budget: L=131072, d=8192, BF16");fig.tight_layout() fig.savefig(OUT/'ring_budget.png',dpi=170) plt.close(fig) print(OUT/'ring_budget.png') 更多 AIGC 论文解读,关注微信公众号「人工智能炼丹君」 每日更新 · 论文精选 · 深度解读 · 技术脉络 微信搜索 人工智能炼丹君 或扫描下方二维码关注
2026年09月05日
10 阅读
0 评论
0 点赞
2026-09-05
AIGC 基本功|张量并行与流水线并行-TP-PP
张量并行与流水线并行 所属方向:分布式训练 | 难度:高阶 | 前置知识:数据并行与 ZeRO 显存切分 关键词:张量并行、流水线并行、Megatron-LM、GPipe、通信开销、并行策略 01. 为什么需要它 ZeRO-3 能把参数、梯度和优化器状态切到数据并行 rank 上,但每一层计算时仍要聚合这一层的参数。若单个 Transformer 层、词表投影或大矩阵乘本身就放不进一张卡,单纯增加数据并行度解决不了问题;若层能放下,但整网常驻状态或激活放不下,也可考虑沿网络深度切开;能否仅靠 ZeRO-3 和重计算解决,需要先算峰值。 这对应两把不同的刀。张量并行(Tensor Parallelism,TP)在一层内部切矩阵,让多张卡共同完成同一个 GEMM;流水线并行(Pipeline Parallelism,PP)沿层切模型,让不同设备各自保存连续或交错的一段层。前者解决“这一层太宽”,后者解决“模型太深”。数据并行沿 batch 复制模型,三者切的是三个不同维度。 先看一个浪费现场:4 个 pipeline stage 处理 1 个 microbatch。第 0 段前向时其余三段都空闲;数据流到第 3 段后,前几段又空闲。若把一个 global batch 切成 8 个 microbatch,让不同样本在各 stage 重叠,理想 forward sweep 从 $4\times8=32$ 个串行 stage-slot 压到 11 个时间槽,但仍有填充和排空的“气泡”。本文脚本算出 4 段、8 个 microbatch 的理想利用率只有 72.73%,气泡占 27.27%。 再看 TP:若一个 MLP 的第一层权重为 $[d,4d]$,可沿输出维切成 $p$ 份,每卡算 $[d,4d/p]$;第二层 $[4d,d]$ 再沿输入维与中间激活匹配切分。这样两次 GEMM 都不需要在中间把完整 $4d$ 激活拼回每张卡,只在必要边界做集合通信。这正是 Megatron-LM 的 column-parallel 与 row-parallel 配对。 难点不在“切成几份”,而在切后保持数学等价,并让通信落在高带宽拓扑上。TP 频繁同步,通常限制在 NVLink/NVSwitch 节点内;PP 在相邻 stage 前向传激活、反向传对应梯度,更适合跨节点。若反过来布置,同一套卡数会被网络延迟拖垮。 02. 最小可用理解 三句话建立框架: TP 切一层的隐藏维、输出通道或 attention head,每层都会通信,换来单层权重、激活和计算分摊。 PP 把层分给不同 stage,再把 batch 切成 microbatch 以重叠各段计算;microbatch 越多,气泡越小,但调度和激活管理更复杂。 生产系统通常组成 DP×TP×PP 的三维网格:TP 放在最快链路内,PP 穿过较慢边界,剩余设备做 DP 与 ZeRO/FSDP。 如果只记一句:TP 是“同一个样本的一层由多卡一起算”,PP 是“同一个样本依次经过多卡上的不同层”。它们都属于模型并行,但通信模式完全不同。 设世界规模 $W$,数据并行度为 $d_p$、张量并行度为 $t_p$、流水线段数为 $p_p$,若没有其他并行轴: $$W=d_p\,t_p\,p_p$$ 这不是越均匀越好。TP 度受矩阵维度、head 数、KV head 数和高速互联限制;PP 度受层数、切分均衡与 microbatch 数限制;最后再用 DP 提升全局吞吐。 03. 数学推导 3.1 Column Parallel:沿输出维切 线性层输入 $X\in\mathbb{R}^{n\times d_{\mathrm{in}}}$,权重 $A\in\mathbb{R}^{d_{\mathrm{in}}\times d_{\mathrm{out}}}$,输出为: $$Y=XA$$ 把 $A$ 沿列切为 $p$ 份: $$A=[A_0,A_1,\ldots,A_{p-1}],\qquad A_i\in\mathbb{R}^{d_{\mathrm{in}}\times d_{\mathrm{out}}/p}$$ 第 $i$ 个 rank 计算: $$Y_i=XA_i$$ 完整结果只是拼接: $$Y=[Y_0,Y_1,\ldots,Y_{p-1}]$$ 前向不必通信,只要输入 $X$ 在 TP 组内复制。若下一层能直接消费分片 $Y_i$,连 all-gather 都可省。反向时,每个 rank 产生一份输入梯度贡献 $dX_i=dY_iA_i^\top$,完整 $dX$ 要求和: $$dX=\sum_{i=0}^{p-1}dX_i$$ 因此 column-parallel 的关键通信落在反向输入梯度 all-reduce。 3.2 Row Parallel:沿输入维切 把输入和权重的匹配维切开: $$X=[X_0,X_1,\ldots,X_{p-1}],\qquad A=\begin{bmatrix}A_0\\A_1\\ \vdots\\A_{p-1}\end{bmatrix}$$ 每卡先算局部部分和 $Z_i=X_iA_i$,完整输出是: $$Y=XA=\sum_{i=0}^{p-1}X_iA_i=\sum_{i=0}^{p-1}Z_i$$ 所以 row-parallel 前向需要 reduce 或 all-reduce,反向则能从复制的 $dY$ 本地得到 $dX_i=dYA_i^\top$。Megatron 把 MLP 的第一层做 column parallel,激活函数逐元素地作用在本地 shard;第二层做 row parallel,最后只求和一次。attention 的 QKV 投影与输出投影也能形成类似配对。 这种设计的精髓不是“每个 Linear 都随便切”,而是让相邻层的分片布局衔接,从而避免在每个算子后都 all-gather。若一个自定义层偷偷需要完整 hidden,通信会突然增加。 3.3 TP 的通信账本 设传输张量有 $N$ 个元素、每元素 $b$ 字节,ring all-reduce 在理想带宽模型中每 rank 通信: $$V_{\mathrm{AR}}=2\frac{p-1}{p}Nb$$ $p$ 是 TP 组大小。Megatron 风格的一个 Transformer 层在前向与反向关键边界各发生集合通信;真实实现可用 reduce-scatter、all-gather 与 sequence parallel 改写,但总成本仍与 token 数、hidden size 和 TP 频率有关。TP 每层同步,所以延迟不能忽略:小矩阵、多层和大 TP 度会让消息切得太碎,算力利用率反而下降。 TP 对参数显存理想降为 $1/p$,但并非所有东西都跟着切。LayerNorm 参数很小,可能复制;embedding 与 LM head 有词表切分规则;激活是否分片取决于 sequence-parallel 配置;通信临时 buffer 还会增加峰值。因此应从真实 module state 和 profiler 测量,不可把全显存直接除以 TP 度。 3.4 PP 气泡从哪里来 设有 $p$ 个等耗时 stage,global batch 切成 $m$ 个 microbatch,每个 stage 处理一个 microbatch 的单向计算耗时为一格。第一个 microbatch 需要 $p$ 格到达尾部,之后每格完成一个。完成一个 forward sweep 共: $$T_{\mathrm{sweep}}=m+p-1$$ 其中每个 stage 真正工作 $m$ 格,所以理想效率与气泡比例: $$\eta_{\mathrm{pipe}}=\frac{m}{m+p-1},\qquad f_{\mathrm{bubble}}=\frac{p-1}{m+p-1}$$ 当 $p=4,m=8$,效率 $8/11=72.73\%$。要到 90% 以上,需要 $m\ge9(p-1)$。这也解释了为什么 stage 很多却只有几个 microbatch 时 PP 很差。 公式假设各段等耗时、通信完全隐藏、无数据依赖停顿。真实训练的气泡还包括 stage 不均衡、点对点传输、参数同步、数据加载与尾 batch。最后一段若有巨大词表投影,前面切得再平均也会被它卡住。 3.5 GPipe、1F1B 与交错调度 GPipe 先把所有 microbatch 前向跑完,再统一反向,调度简单但要保存更多未反向的激活。1F1B 在 warmup 后交替执行一次 forward 和一次 backward,使稳态内存更接近少量 microbatch;同步更新语义仍要求一个 global batch 的梯度累积完再 step。 交错 1F1B 让每个物理设备承载多个 virtual stage,把一个大气泡拆细。它可改善负载与气泡,却增加更多通信边界、依赖和调度复杂度。Pipeline 并行的优化对象不只是公式里的 bubble,还包括峰值激活、通信重叠和 kernel 连续性。 均衡同步流水线的理想利用率,按 m/(m+p−1) 计算。增加 microbatch 能摊薄填充与排空气泡;图中不含通信、stage 失衡与小 GEMM 效率变化。 04. 代码实现 tp_linear_sim.py 只用标准库,构造 $X:[2,4]$、$W_1:[4,6]$、$W_2:[6,3]$。先完整计算,再把第一层按列、第二层按行切到两个虚拟 rank: X shape=(2, 4), W1 shape=(4, 6), W2 shape=(6, 3), tp=2 column shards: W1=(4, 3) each, hidden=(2, 3) each row shards: W2=(3, 3) each, partial output=(2, 3) each full output=[[52, 42, 72], [132, 114, 180]] max hidden diff=0, max output diff=0 输出差为 0,验证“按列拼接、按行求和”与完整矩阵乘严格等价。真实浮点并行的归约顺序不同,通常会出现末位误差,不能要求 bitwise 一致,而应设置与 dtype 相称的容差。 pipeline_bubble.py 默认模拟 4 段、8 个 microbatch: stages=4, microbatches=8, ideal_slots_per_sweep=11 pipeline_efficiency=72.73%, bubble_fraction=27.27% forward schedule (slot -> active stage:microbatch) 00: S0:M0 01: S0:M1 S1:M0 02: S0:M2 S1:M1 S2:M0 03: S0:M3 S1:M2 S2:M1 S3:M0 04: S0:M4 S1:M3 S2:M2 S3:M1 05: S0:M5 S1:M4 S2:M3 S3:M2 06: S0:M6 S1:M5 S2:M4 S3:M3 07: S0:M7 S1:M6 S2:M5 S3:M4 08: S1:M7 S2:M6 S3:M5 09: S2:M7 S3:M6 10: S3:M7 最前面三格在填充,最后三格在排空,中间四个 stage 才全部忙碌。用参数改变 stages 与 microbatches,就能看到增加 microbatch 如何摊薄固定气泡。 05. 工业级实现对照 以 2026-09 为准,NVIDIA Megatron-LM 的 tensor_parallel/layers.py 仍提供 ColumnParallelLinear 与 RowParallelLinear。工业实现除切权重外,还管理参数初始化、bias、是否 gather 输出、异步梯度 all-reduce、梯度累积融合、sequence parallel、专家并行组和通信 buffer。最小脚本只证明线性代数,没有模拟 autograd 与 NCCL。 ColumnParallelLinear 的 gather_output 决定输出是否立刻汇总;RowParallelLinear 的 input_is_parallel 表示输入是否已经按最后一维切好。错误组合往往不是数值报错,而是多做一次 gather 或得到错位 shard。审计并行模型时,应给每个边界标清 global shape、local shape、分片轴、复制轴和预期 collective。 PP 代码位于 Megatron Core 的 pipeline_parallel 调度模块。配置中的 pipeline_model_parallel_size 切物理 stage,virtual_pipeline_model_parallel_size 启用交错。系统必须传递前向激活和反向梯度,处理 tied embedding、首尾 stage 特殊 loss、不同张量 shape、激活释放与通信重叠。 Megatron 2021 展示了 TP、PP 与 DP 的组合,并提出交错流水调度。工程上的常见布局是:同一节点内组成 TP 组,邻接节点组成 PP 链,跨副本组成 DP 组。原因是 TP 每层通信而 PP 只在 stage 边界传输;把频繁 collective 放在更快互联上更划算。 生产配置还要做 layer partition。按层数平均只在每层耗时相同才合理;MoE 层、cross-attention、视觉模块、embedding 与词表 head 的成本差异很大。应先 profile 单层前后向时间和激活尺寸,再按时间而不是按层数切 stage。 06. 代价与边界 TP 的主要代价是高频 collective。TP 度增加后,每卡 GEMM 变小,计算效率下降,通信延迟占比上升;跨节点 TP 尤其敏感。维度还必须能合理整除:attention head、KV head、MLP intermediate size、词表分片与低精度 tile 对齐都会形成约束。 PP 的主要代价是气泡、激活驻留与调度复杂度。microbatch 增多能降气泡,却让每次 GEMM 的 batch 更小,可能降低算力利用率;还会增加调度次数。梯度累积数也受 global batch、DP 度与 microbatch size 约束: $$B_{\mathrm{global}}=B_{\mathrm{micro}}\,m\,d_p$$ 为了把 $m$ 调大而偷偷改变 global batch,会连学习率与收敛语义一起改变。更稳妥的做法是减小 microbatch size、保持 global batch,再检查小 GEMM 是否仍高效。 PP stage 之间有顺序依赖,单个慢 stage 会拖住全链。设备故障、动态 shape、条件分支和 MoE 路由也比普通 DP 难处理。推理时 batch 与请求长度动态变化,训练得到的均衡切分未必仍均衡。 TP 与 ZeRO/FSDP 并非天然可任意叠加。两者可能切同一参数的不同轴,process group、参数初始化、checkpoint layout 和 optimizer state 都需协调。首先建立二维或三维 rank 映射,再启用框架明确支持的组合,避免自行套两层 wrapper。 什么时候不该用?模型单卡能放下且吞吐受数据不足或 CPU 限制时,TP/PP 只会增加同步;单层很小而层很多时,优先 PP 或 FSDP;单层巨大而层数不多时,TP 更直接;超长序列导致激活爆炸时,还要引入下一篇的序列/上下文并行。 6.1 从模型形状反推 TP 度 假设 decoder hidden size 为 $d$,MLP expansion 为 $4d$,attention head 数为 $H$,每头维度为 $d_h=d/H$。采用 $t_p$ 路 TP 时,至少希望 $H/t_p$ 与 $4d/t_p$ 为整数,且本地矩阵维度满足 Tensor Core tile 对齐。GQA 还要检查 KV head 数 $H_{KV}$;若 Q head 可整除而 KV head 不可整除,框架可能复制 K/V 或根本拒绝配置。 参数容量只是下界。以 MLP 两个权重为例,忽略 bias: $$P_{\mathrm{MLP}}=d(4d)+(4d)d=8d^2$$ 理想 TP 后每卡为 $8d^2/t_p$ 个参数。然而输入 $X$、残差、LayerNorm 与某些输出仍可能复制。是否启用 sequence parallel,会决定 $[B,L,d]$ 激活是每 TP rank 一份还是沿 token 切开。评估 TP 度时要分别列参数、可分片激活、复制激活和通信 buffer,不能只用总显存除 $t_p$。 性能上,本地 GEMM 的算术强度会随 $t_p$ 增大而下降。若 $d/t_p$ 太小,矩阵乘无法占满 SM,即使通信为零也会变慢。实际选型一般从“能放下的最小 TP 度”起步,再尝试少数相邻值;不是卡越多 TP 越大。 6.2 三维 rank 映射为什么决定速度 有 2 个节点、每节点 8 卡,总计 16 卡。假设 TP=8、PP=2、DP=1,合理映射是每个节点内部构成一个 TP 组,两个节点作为相邻 PP stage。这样每层 TP collective 走 NVLink/NVSwitch,只有 stage 边界激活跨节点。 若 rank 编号错误,使每个 TP 组横跨两节点,那么每个 Transformer 层的 all-reduce 都经过网络;PP 反而在节点内。数学结果完全正确,吞吐却可能大幅下降。这类问题从配置数字看不出来,必须导出每个 process group 的物理 GPU、主机名、PCIe/NVLink 路径和 NIC 亲和性。 当 DP>1 时,DP 组应从相同 TP/PP 坐标上取不同副本。例如把 rank 写成坐标 $(r_d,r_p,r_t)$,模型同一 shard 的梯度只在 $r_d$ 维同步;层内 tensor collective 只改变 $r_t$;pipeline 点对点只沿 $r_p$ 邻接。把坐标语义明确下来,checkpoint 分片和故障定位才不会依赖偶然 rank 编号。 多 NIC 节点还要关注通信并发。TP all-reduce、PP send/recv、DP reduce-scatter 可能同时争同一链路。单独 benchmark 每个 collective 很快,不代表组合后仍快;应在真实 schedule 上看每条 stream 和 NIC 的时间线。 6.3 Pipeline 内存不是简单除以段数 PP 将参数按层分段,参数内存近似降为 $1/p_p$,但激活取决于 schedule。GPipe 前向完所有 $m$ 个 microbatch 才反向,每 stage 可能保留 $O(m)$ 份边界与层内激活。1F1B warmup 后尽早反向,可显著减少同时存活的 microbatch 数;不同 stage 的 warmup 长度又不同,峰值不完全一致。 activation checkpoint 将层内保存换成反向重算,但 stage 边界张量通常仍要保留或重新通信。若 PP 与 ZeRO-3 叠加,重算还可能再次触发参数 all-gather;如果 prefetch 与 release 时机不匹配,理论显存节省会被临时完整参数覆盖。 一个更实用的峰值模型是: $$M_s=M_{\mathrm{params},s}+n_{\mathrm{live},s}M_{\mathrm{act/micro},s}+M_{\mathrm{comm},s}+M_{\mathrm{workspace},s}$$ $s$ 是 stage,$n_{\mathrm{live},s}$ 是该调度下同时存活的 microbatch 数。选切分点的目标应是最小化 $\max_s M_s$ 并平衡每段时间,而不是让每段层数相同。首段 embedding、尾段 vocab projection 和 loss 往往需要单独计量。 6.4 Schedule 的一致性与权重版本 同步 GPipe 或 1F1B 中,一个 global batch 的所有 microbatch 应使用同一版参数,梯度累积后再统一 optimizer step。因此它与非流水同步训练在数学上等价,只改变操作顺序。随机 dropout 若要严格比较,还需保证不同 schedule 消耗一致的随机数流。 异步 pipeline 为减少 flush 可能允许某些 microbatch 使用旧权重,产生 weight staleness。PipeDream 的权重暂存与调度就是为管理这个问题。吞吐更高不等于优化轨迹相同;对需要可复现或大规模预训练的任务,同步 schedule 通常更容易验收。 梯度累积的 loss normalization 也常出错。若每个 microbatch loss 已取 mean,再把 $m$ 份梯度直接相加,最终梯度比全局 mean 大 $m$ 倍;框架可能在 loss、backward 或 optimizer wrapper 某处除 $m$。迁移实现时必须确认缩放发生在哪里,特别是最后一个不完整 microbatch。 6.5 如何系统调优而不是枚举所有组合 第一步做容量约束:根据参数、优化器、激活和临时 buffer,排除会 OOM 的 TP/PP。第二步做整除约束:检查层数、head、KV head、MLP width、词表与 microbatch。第三步按拓扑放组:TP 留在最快域,PP 穿过节点,DP 使用余下副本。第四步才短跑候选方案。 短跑至少记录每个 stage 的 forward/backward 时间、TP collective、PP send/recv、DP collective、bubble、MFU 和峰值显存。若 stage 时间方差大,先重切层;若 collective 裸露,检查异步与 bucket;若 GEMM 利用率低,减少 TP 或增大 microbatch;若 bubble 大,增加 microbatch 或 virtual stage。 还要避免只优化稳态。训练中 checkpoint、评估、数据切换、动态 loss scale 与长短样本混合会改变节拍。视频模型的序列长度可能随分辨率和帧数变化,固定层切分在不同 batch 上会失衡。可以按长度分桶、限制每批 token 数,或使用能处理动态 shape 的 schedule,但都要重新验证 global batch 语义。 6.6 Checkpoint 与并行度变更 TP checkpoint 的一个权重可能沿行或列分片,PP checkpoint 又只在拥有该层的 stage 上出现。保存时应记录全局 shape、分片轴、offset、replica group 与 tied-weight 关系。仅按 rank 存文件却没有布局元数据,会显著增加从 TP=8 改成 TP=4 的恢复难度;若完整掌握原切分约定仍可转换,但必须验证。 成熟 distributed checkpoint 会把逻辑参数名与物理 shard 解耦,加载时重新规划切片。验证转换不能只看“文件读完”,应抽样 all-gather 后与原始全参数比较,并跑一个确定性 forward。optimizer state 也必须按相同参数布局重分片,特别是 Adam 的 m、v 与 FP32 主权重。 PP 的 tied embedding 是典型边界:输入 embedding 在首段,输出投影在尾段,但两者可能共享参数。系统要么在两个 stage 间同步梯度,要么采用明确的复制与更新协议。忽略它会让模型能跑、loss 也下降,却不再是原架构。 最后为每个组合保留机器可读配置:world size、各轴度数、rank 坐标映射、global/micro batch、累积数、schedule、virtual stage、precision 与 checkpoint schema。没有这份清单,性能回归时很难判断是代码变化还是并行布局变化。 6.7 一个具体的 64 卡选型例子 假设 8 个节点、每节点 8 卡,模型单层在 4 卡上能放下,96 层在单节点放不下,目标 global batch 又允许 4 个数据副本。可先试 TP=4、PP=4、DP=4,乘积正好 64。每个节点容纳两个 TP 组,相邻两个节点组成一条四段 pipeline;相同 TP/PP 坐标跨四个副本形成 DP 组。 为什么不直接 TP=8、PP=2、DP=4?它减少 PP 气泡,但本地 GEMM 变小、每层 collective 参与卡数翻倍。为什么不 TP=2、PP=8、DP=4?单层可能放不下,而且 PP 段更多,要求更多 microbatch 才能摊薄气泡。三个方案都满足乘积约束,只有容量、profile 和拓扑能决定赢家。 设 PP=4、microbatch 数 $m=16$,理想气泡为 $3/19=15.79\%$;若为配合 DP=4,global batch 满足 $B_{\mathrm{global}}=B_{\mathrm{micro}}\times16\times4$。当目标 global batch 固定时,microbatch size 可能被压得太小。此时交错 PP 可在不继续增加 $m$ 的情况下缩小气泡,但要多传 stage 边界。 选定后再检查每卡峰值。若尾段词表 head 使 stage 3 明显更慢,可把少量 Transformer 层从尾段移到前段;若 stage 0 embedding 占显存而计算很少,切分目标应同时满足容量与时间,而不是追求参数量绝对相等。 6.8 故障现象与定位路径 若所有 GPU 利用率呈周期性锯齿且空白集中在迭代首尾,优先看 PP bubble;若每层 GEMM 后都有长 NCCL 条带,优先看 TP 通信或 rank 跨节点;若只有某一个 stage 长期满载、其他 stage 等待,是层切分失衡;若迭代末尾集中等待,是 DP 梯度同步没有充分 overlap。 出现 hang 时,先核对各 rank collective 调用序列。条件分支导致某些 rank 少调用一次 all-reduce,或 PP 两端 send/recv shape 不一致,都会永久等待。为通信操作记录 group、sequence number、peer、shape 与 dtype,比只看 Python stack 更有用。设置超时只能让错误更快暴露,不能修复顺序。 数值不一致时,从一个微型模型开始,关闭 dropout 和 fused kernel,分别测试 TP、PP,再测试组合。TP 常见错误是切分轴、bias 重复相加、输出误 gather;PP 常见错误是 loss 缩放、跨段激活 requires-grad、共享参数同步。一次只启用一根并行轴,能把搜索空间从三维降为一维。 OOM 若只发生在第一步,往往是 optimizer state 首次建立或通信 bucket 惰性分配;若在若干步后发生,可能是 graph retention、动态 shape 缓存或碎片;若只在某个 PP stage,先看该段真实 activation 与临时 buffer。不要看到 OOM 就统一减 microbatch,这可能掩盖泄漏却降低所有卡效率。 6.9 上线验收清单 发布配置前确认:世界规模等于各并行轴乘积;每个 rank 坐标唯一;TP group 位于预期高速域;head、KV head、MLP 与词表可整除;各 stage 时间和峰值接近;global batch 算术正确;loss normalization 不依赖 microbatch 数;tied weight 有同步协议;checkpoint 能跨目标并行度恢复。 性能门禁同时保存 tokens/s、MFU、bubble、每类 collective 的裸露时间、峰值显存和最慢 stage。只保存总迭代时间无法解释回归。集群拓扑或通信库升级后,即使模型代码未变,也要重跑基准,因为并行策略本质上是对物理系统的映射。 最后预留降级方案:某节点高速链路异常时能否减少 TP、增加 PP;某个 checkpoint 是否可重分片;global batch 是否仍保持;恢复后数值是否连续。真正稳健的并行配置,不只是峰值最快,也要能在硬件变化和断点恢复时保持语义清楚。 6.10 理论通信量为什么不等于训练时间 同样传 1 GiB,一个大 all-reduce 与数百个小 collective 的耗时不同。常用模型 $T=\alpha n+\beta V$ 中,$n$ 是消息次数,$V$ 是字节量,$\alpha$ 表示每次启动延迟,$\beta$ 表示每字节时间。TP 在每层频繁通信,尤其受 $\alpha$ 影响;PP 消息次数少,但单次边界激活可能很大。 通信是否裸露还取决于依赖。梯度 all-reduce 可与更早层 backward 重叠,PP send 可与下一份本地计算重叠;位于关键路径上的 collective 即使字节少也会直接延长 step。性能报告应区分总 NCCL 时间与 exposed communication time。 bucket 太小会增加消息次数,太大又推迟通信启动,减少 overlap。最优 bucket 与层大小、网络延迟和反向节奏有关。框架默认值是通用折中,不一定适合视频 DiT 的大激活或 MoE 的不均匀参数。 最后还要观察尾延迟。一个 rank 因热降频、ECC、数据抖动或网络拥塞变慢,collective 会让整个组等待。平均 GPU 时间看似健康,最慢 rank 才决定训练。按 rank 记录分位数,并把异常节点与拓扑关联,是大规模并行调优的基本动作。 理论模型最终要由 trace 校准。为每次实验保存并行配置、网络拓扑、逐 stage 时间线与通信矩阵,下一次才能定位回归。若只留一行 tokens/s,既无法判断瓶颈是 GEMM、气泡还是慢 rank,也无法安全迁移到另一代 GPU。并行策略不是模型的附属启动参数,而是计算图和硬件共同组成的一部分,理应像模型结构一样接受版本管理和回归测试。 07. 经典论文脉络 GPipe(arXiv:1811.06965):用 microbatch 流水化跨加速器的模型分段,并以同步 mini-batch 语义训练巨型网络。 PipeDream(arXiv:1806.03377):探索 1F1B 与异步流水,揭示吞吐、权重版本和一致性之间的权衡。 Megatron-LM(arXiv:1909.08053):提出 Transformer 内高效的 tensor model parallel 切法,用少量 collective 支撑数十亿参数。 Efficient Large-Scale Language Model Training(arXiv:2104.04473):系统组合 TP、PP、DP,并以交错流水减少气泡,扩展到千卡与万亿参数。 论文演进说明:单一并行轴只能解决一种容量瓶颈,规模继续增长后,真正的问题变成如何把多个轴映射到硬件拓扑并共同调度。 08. 常见误解 误解一:TP 就是把每层平均切开。 切分轴决定通信。Megatron 的 column-row 配对是为了让中间分片直接衔接,不是机械地切一半权重。 误解二:PP 段数越多,显存越省且速度越快。 参数容量会下降,但固定气泡随 $p-1$ 增长;若 microbatch 不够,更多 stage 反而更闲。 误解三:microbatch 越多越好。 它降低理想气泡,却减小 GEMM、增加调度,并可能扩大激活队列。要联动测 MFU 与显存。 误解四:1F1B 是异步优化。 常用同步 1F1B 只是改变前后向顺序,仍在一个 global batch 梯度完成后统一更新;PipeDream 式异步才涉及权重陈旧。 误解五:TP 能把所有显存除以 TP 度。 小参数、复制激活、通信 buffer 与 runtime 不会理想均分。必须看逐项账本。 误解六:三维并行度乘起来等于卡数就配置正确。 整除只是必要条件。拓扑、矩阵 tile、head 数、stage 均衡、global batch 都可能让配置不可用。 09. 动手验证 先运行 tp_linear_sim.py,把 TP 从 2 改成能整除维度的 3,重新切 $W_1$ 的输出与 $W_2$ 的输入。预期完整输出仍一致。然后故意让第二层 shard 顺序交换,结果会静默错误;这说明 distributed tensor layout 必须携带明确语义。 运行 pipeline_bubble.py,固定 stages=8,分别取 microbatches=1、8、32、72。理论效率依次为 $1/8$、$8/15$、$32/39$、$72/79$。脚本当前只实现均衡 stage 的闭式公式;扩展成事件调度模拟后,再给每个 stage 设置不同时间,以最慢 stage 为节拍比较,观察均分层数为何不等于均分时间。 在真实集群做 TP=1/2/4/8 对照,固定 global batch 和模型,记录每层 GEMM 时间、collective 时间、MFU、峰值显存。预期早期 TP 能解除容量或提高总吞吐,超过某点后小 GEMM 和通信使收益反转。用拓扑工具确认 TP rank 是否真的落在 NVLink 域内。 对 PP 做相同实验:固定 PP 度,逐步增加 microbatch 数;再固定 microbatch,增加 virtual stage。记录 bubble、激活峰值与端到端 tokens/s。不要只看 framework 打印的理论 bubble,GPU timeline 中的空白才是真实气泡。 最后做一次等价性验收:单卡、TP、PP、TP+PP 使用相同初始化和样本,关闭 dropout,比较一个 step 的 loss 和若干参数梯度。低精度与归约顺序会造成小差异,但若误差随层爆炸,优先检查 shard 顺序、bias、loss 归一化与 tied weight。 10. 延伸阅读 读完这篇可以继续看: 数据并行与 ZeRO 显存切分:DP 沿 batch 复制计算、ZeRO 沿 DP 组切状态,是三维并行的第一根轴(已发布)。 序列并行与 Ring Attention:当瓶颈从模型宽度转向超长 token,沿 sequence 切激活与 attention(本文系列下一篇)。 混合精度与数值稳定性:低精度改变 TP/PP 通信字节与归约误差,配置并行前要先明确累加精度(本文系列上一篇)。 选型时先问容量瓶颈在哪里:单层太宽选 TP,层总数太深选 PP,副本吞吐选 DP,超长激活选 SP。随后才是把这些轴映射到实际互联。公式给出上限,profile 决定最终配置。 附录:完整代码 09 节用到的脚本全文如下(tp_linear_sim.py、pipeline_bubble.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 tp_linear_sim.py #!/usr/bin/env python3 """Verify column- and row-parallel linear algebra without any framework.""" def matmul(a, b): return [[sum(x * y for x, y in zip(row, col)) for col in zip(*b)] for row in a] def split_columns(a, parts): width = len(a[0]) // parts return [[row[i * width:(i + 1) * width] for row in a] for i in range(parts)] def split_rows(a, parts): height = len(a) // parts return [a[i * height:(i + 1) * height] for i in range(parts)] def add(a, b): return [[x + y for x, y in zip(ra, rb)] for ra, rb in zip(a, b)] x = [[1, 2, 3, 4], [5, 6, 7, 8]] # [tokens=2, hidden=4] w1 = [[1, 0, 2, 0, 3, 0], [0, 1, 0, 2, 0, 3], [1, 1, 1, 1, 1, 1], [2, 1, 0, 1, 2, 1]] # [4, 6] w2 = [[1, 0, 1], [0, 1, 1], [1, 1, 0], [2, 0, 1], [0, 2, 1], [1, 0, 2]] # [6, 3] full_hidden = matmul(x, w1) full_output = matmul(full_hidden, w2) # Column parallel W1: every rank computes a slice of output features. w1_shards = split_columns(w1, 2) hidden_shards = [matmul(x, shard) for shard in w1_shards] column_joined = [left + right for left, right in zip(*hidden_shards)] # Row parallel W2: input features and W2 rows use matching shards, then sum. w2_shards = split_rows(w2, 2) partials = [matmul(h, w) for h, w in zip(hidden_shards, w2_shards)] row_reduced = add(partials[0], partials[1]) max_hidden_diff = max(abs(a - b) for ra, rb in zip(full_hidden, column_joined) for a, b in zip(ra, rb)) max_output_diff = max(abs(a - b) for ra, rb in zip(full_output, row_reduced) for a, b in zip(ra, rb)) print("X shape=(2, 4), W1 shape=(4, 6), W2 shape=(6, 3), tp=2") print("column shards: W1=(4, 3) each, hidden=(2, 3) each") print("row shards: W2=(3, 3) each, partial output=(2, 3) each") print(f"full output={full_output}") print(f"max hidden diff={max_hidden_diff}, max output diff={max_output_diff}") pipeline_bubble.py #!/usr/bin/env python3 """Idealized GPipe bubble calculator; stdlib only.""" import argparse parser = argparse.ArgumentParser() parser.add_argument("--stages", type=int, default=4) parser.add_argument("--microbatches", type=int, default=8) args = parser.parse_args() if args.stages < 1 or args.microbatches < 1: parser.error("stages and microbatches must be positive") p, m = args.stages, args.microbatches slots = m + p - 1 efficiency = m / slots print(f"stages={p}, microbatches={m}, ideal_slots_per_sweep={slots}") print(f"pipeline_efficiency={efficiency:.2%}, bubble_fraction={1-efficiency:.2%}") print("forward schedule (slot -> active stage:microbatch)") for t in range(slots): active = [f"S{s}:M{t-s}" for s in range(p) if 0 <= t - s < m] print(f"{t:02d}: " + " ".join(active)) make_figures.py """Regenerate this article's deterministic teaching figure (numpy + matplotlib).""" from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np OUT=Path(__file__).resolve().parents[1]/"figures" OUT.mkdir(exist_ok=True) plt.rcParams.update({"font.sans-serif":["PingFang SC","Arial Unicode MS","DejaVu Sans"],"axes.unicode_minus":False}) m=np.arange(1,129) fig,ax=plt.subplots(figsize=(8,4.8)) for p in [2,4,8,16]:ax.plot(m,m/(m+p-1),label=f"{p} stages") ax.set(xlabel="Microbatches per global batch",ylabel="Ideal utilization",ylim=(0,1.02),title="Balanced synchronous pipeline: m / (m + p - 1)") ax.legend();ax.grid(alpha=.25);fig.tight_layout() fig.savefig(OUT/'pipeline_utilization.png',dpi=170) plt.close(fig) print(OUT/'pipeline_utilization.png') 更多 AIGC 论文解读,关注微信公众号「人工智能炼丹君」 每日更新 · 论文精选 · 深度解读 · 技术脉络 微信搜索 人工智能炼丹君 或扫描下方二维码关注
2026年09月05日
5 阅读
0 评论
0 点赞
2026-09-05
AIGC 基本功|混合精度与数值稳定性-AMP
混合精度与数值稳定性 所属方向:分布式训练 | 难度:进阶 | 前置知识:无 关键词:混合精度、FP16、BF16、FP8、loss scaling、溢出、数值稳定 01. 为什么需要它 同一份训练代码,把 BF16 改成 FP16,吞吐可能更高,也可能几十步后 loss 直接变成 NaN。另一个常见现场是:模型看起来正常收敛,但某些参数长期没有变化;检查梯度才发现,一些绝对值不超过 $2^{-25}\approx2.98\times10^{-8}$ 的梯度在 round-to-nearest-even 写回 FP16 时已经变成了 0。 这不是“半精度不准”一句话能解释的。训练里的风险至少有三种:数太大,超过格式的最大有限值而溢出成 inf;数太小,舍入后落到 0(或硬件将次正规数 flush-to-zero);数虽然在范围内,却因为有效位太少,在加法里被大数吞掉。三者的修复手段不同:loss scaling 能救小梯度,却救不了激活溢出;换 BF16 能扩展动态范围,却不会自动改善尾数精度;把归约留在 FP32 能减小累加误差,却不代表所有输入和权重都要回到 FP32。 为什么还要承担这些麻烦?因为 Transformer 的大头通常是矩阵乘。现代加速器对低精度矩阵乘提供更高吞吐,权重和激活每元素从 4 字节降到 2 字节也能减小显存与带宽压力。Mixed Precision Training 给出的核心配方是:低精度做大部分前后向,保留 FP32 主权重进行更新,并对 loss 缩放以防小梯度消失;其模型显存可接近减半,同时保持精度。 “混合”的重点不是选一个统一 dtype,而是给不同数值角色分工。大 GEMM 适合低精度输入,softmax、归一化、loss 和长归约通常需要更大的范围或更高精度,优化器状态又有自己的要求。AMP 的价值正是把逐算子的 dtype 路由表和梯度缩放流程标准化。 一个数字能说明 FP16 与 BF16 的性格差异:FP16 最大有限值只有 65504,而 BF16 与 FP32 一样有 8 位指数,最大值约 $3.39\times10^{38}$。反过来,FP16 有 10 位小数尾数,1 附近的间隔约 $9.77\times10^{-4}$;BF16 只有 7 位尾数,间隔约 $7.81\times10^{-3}$。BF16 更不容易炸,FP16 在可表示范围内更细;“更稳定”与“更精确”不是同一维度。 02. 最小可用理解 三句话先建立框架: AMP 让适合 Tensor Core 的矩阵乘使用 FP16/BF16,让 softmax、归约等敏感算子保留 FP32,而不是粗暴地把整个模型全部转成 half。 FP16 训练通常用 loss scale $S$ 把反向梯度整体放大,更新前再除回去;发现 inf/NaN 时跳过这一步并缩小 $S$。 BF16 动态范围接近 FP32,通常不需要 GradScaler,但尾数更短;FP8 则必须进一步管理张量尺度、绝对最大值历史和累加精度。 如果只记一个检查顺序:先问异常发生在前向还是反向,再区分 overflow、underflow 与 rounding,最后才决定换 dtype、调 scale,还是把局部算子提升到 FP32。 AMP 中常见的四种角色也要分开:参数存储 dtype、算子输入 dtype、乘法累加 dtype、优化器状态 dtype。日志打印“BF16 training”并不能证明四者都是 BF16。很多 GEMM 是 BF16 输入、FP32 累加;Adam 的一阶矩和二阶矩仍为 FP32;某些框架还会保存 FP32 主权重。 03. 数学推导 3.1 浮点数到底牺牲了什么 对非零正规数,一个二进制浮点数可写成(次正规数没有隐含的前导 1): $$x=(-1)^s(1.f)_2\,2^{e-\mathrm{bias}}$$ $s$ 是符号位,$e$ 是指数域,$f$ 是小数域。指数位数决定动态范围,尾数位数决定相邻可表示数的间距。FP16 是 1 位符号、5 位指数、10 位小数;BF16 是 1、8、7;FP32 是 1、8、23。于是: FP16:最小正规数 $2^{-14}=6.1035\times10^{-5}$,最小次正规数 $2^{-24}=5.9605\times10^{-8}$,最大有限值 65504; BF16:最小正规数 $2^{-126}\approx1.1755\times10^{-38}$,最大有限值约 $3.3895\times10^{38}$; FP32:动态范围与 BF16 同阶,但 23 位小数显著降低舍入误差。 机器 epsilon 是 1 与下一个可表示数的间隔;下面列的是 epsilon。在 round-to-nearest 模式下,通常定义的 unit roundoff 为这些值的一半: $$\epsilon_{\mathrm{FP16}}=2^{-10},\qquad \epsilon_{\mathrm{BF16}}=2^{-7},\qquad \epsilon_{\mathrm{FP32}}=2^{-23}$$ 因此 $1+10^{-4}$ 写入 FP16 或 BF16 都可能仍是 1。更危险的是参数更新:若权重 $w=1$,学习率乘梯度只有 $10^{-5}$,直接在低精度权重上做 $w-\eta g$,变化会被舍掉。FP32 主权重就是在高精度副本上累计细小更新,再把结果舍入给低精度前后向。 在 round-to-nearest-even 且保留次正规数时,绝对值不超过 $2^{-25}$ 的 FP16 输入舍入为零;介于这个阈值和 $2^{-24}$ 之间的数可能舍入为最小次正规数,而非一律为零。支持 flush-to-zero 的执行路径还需另查硬件与算子规则。 3.2 为什么 loss scaling 不改梯度 设损失为 $L(\theta)$,参数为 $\theta$,真实梯度为 $g=\nabla_\theta L$。把 loss 乘常数 $S$ 后反向: $$g_s=\nabla_\theta(SL)=S\nabla_\theta L=Sg$$ 只要在优化器读取梯度前除以 $S$,就恢复原梯度: $$g=\frac{g_s}{S}$$ 关键是量化顺序。若 $g=10^{-8}$,先写入 FP16 会变 0,之后再乘任何数都救不回来;若先由链式法则得到 $Sg=1.024\times10^{-5}$,它能被 FP16 表示,写回后再用 FP32 除以 1024,就能保留接近 $10^{-8}$ 的非零值。 静态 scale 要人工选 $S$。太小救不了下溢,太大又会让大梯度超过 65504。动态 GradScaler 维护随训练变化的 $S_t$: $$S_{t+1}=\begin{cases}\beta S_t,&g_s\text{ 含 inf/NaN}\\ \gamma S_t,&\text{连续 }K\text{ 步有限}\\ S_t,&\text{其他情况}\end{cases}$$ $\beta<1$ 是回退因子,$\gamma>1$ 是增长因子,$K$ 是增长间隔。出现非有限梯度时必须跳过 optimizer step,否则坏更新会污染权重和 Adam 状态。PyTorch 文档还提醒:scale 不保证始终大于 1,BF16 预训练模型强转 FP16 时可能因数值过大而一路回退。 3.3 为什么 softmax 和归约容易出事 直接计算 softmax: $$p_i=\frac{e^{x_i}}{\sum_j e^{x_j}}$$ 若 $x_i$ 很大,指数会溢出。稳定写法先减最大值 $m=\max_jx_j$: $$p_i=\frac{e^{x_i-m}}{\sum_j e^{x_j-m}}$$ 这样最大指数为 1,不改变结果,却压下溢出风险。LayerNorm 方差也不宜用 $E[x^2]-E[x]^2$ 的低精度朴素形式:两个接近的大数相减会灾难性消减。长向量求和误差还会随项数累积,因此 AMP 通常让归约在 FP32 完成。 还要区分“算子输出 dtype”和“内部累加 dtype”。矩阵乘输入、输出可以是 BF16,但硬件乘积常进入 FP32 accumulator,最终再舍入回 BF16。若自定义 kernel 把 accumulator 也降成 16 位,它不再与常规 AMP 拥有同样数值性质。 3.4 FP8 为什么不只是再少 8 位 FP8 常见 E4M3 与 E5M2。FP8 Formats for Deep Learning 用 E4M3 提供较多有效位,用 E5M2 提供较大范围。原始张量很难恰好落进狭窄范围,通常给每个张量或块维护尺度 $a$: $$q=Q_{\mathrm{FP8}}(x/a),\qquad x_{\mathrm{deq}}=a\,q$$ 尺度可由绝对最大值 amax 与格式上限估计。若每步都同步全局 amax,又会引入开销;工业实现会用历史窗口、延迟缩放或块缩放。FP8 的正确性同时依赖格式、粒度、尺度更新、异常值分布和高精度累加,不能机械替换 dtype 字符串。 直接量化的相对误差。小数值端 FP16 可能舍入为零(相对误差 1),超过最大有限值后溢出;BF16 的范围更宽,但正规数的尾数精度更低。精确可表示点的误差为零,图中作绘图下限处理。 04. 代码实现 两个脚本都只依赖 Python 标准库。float_formats.py 用 IEEE 二进制打包模拟 FP16,并先按舍入位加偏置,再截去 FP32 低 16 位,实现 BF16 round-to-nearest-even。真实输出: format exp frac min_normal max_finite epsilon_at_1 FP16 5 10 6.1035e-05 6.5504e+04 9.7656e-04 BF16 8 7 1.1755e-38 3.3895e+38 7.8125e-03 FP32 8 23 1.1755e-38 3.4028e+38 1.1921e-07 value -> FP16 | BF16 1.0001 -> 1 | 1 1e-05 -> 1.00136e-05 | 1.00136e-05 100000 -> inf | 99840 100000 在 FP16 变成 inf,在 BF16 仍有限;1.0001 在两种 16 位格式里都舍入成 1。动态范围与精度的区别由此可见。 loss_scaling_sim.py 把五个小梯度直接量化,再先乘 1024、量化、最后除回去: shape=(5,), loss_scale=1024 gradient direct_fp16 scaled_fp16/unscaled 1.000e-08 0.000e+00 1.001e-08 3.000e-08 5.960e-08 2.998e-08 1.000e-07 1.192e-07 1.000e-07 1.000e-06 1.013e-06 1.000e-06 1.000e-05 1.001e-05 9.999e-06 nonzero: direct=4/5, scaled=5/5 $10^{-8}$ 直接写 FP16 已为 0,缩放路径却保留为 $1.001\times10^{-8}$。这不是凭空增加精度:量化误差仍在,只是把数搬进可表示区间。 真实 PyTorch CUDA 训练的核心顺序: scaler = torch.amp.GradScaler("cuda") for inputs, targets in loader: optimizer.zero_grad(set_to_none=True) with torch.autocast("cuda", dtype=torch.float16): loss = loss_fn(model(inputs), targets) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) scaler.step(optimizer) scaler.update() 梯度裁剪必须在 unscale 之后,若对放大后的梯度仍用阈值 C 裁剪,反缩放后的有效阈值会变成 C/S。若用 BF16,通常保留 autocast 而不用 GradScaler;是否省略仍应由实际梯度分布验证。 05. 工业级实现对照 以 2026-09 的实现为准,PyTorch 的 torch/amp/grad_scaler.py 中 GradScaler.scale 把 loss 乘当前尺度;unscale 按 device 与 dtype 分组检查并反缩放梯度;step 只在没有 inf/NaN 时调用优化器;update 再根据各 optimizer 的 found-inf 状态调整尺度。它还处理稀疏梯度、多设备、多优化器、checkpoint state dict 和惰性初始化,这些都是最小模拟没有覆盖的工程边界。 官方 AMP 文档 当前推荐统一使用 torch.autocast 与 torch.amp.GradScaler;旧的 torch.cuda.amp 接口已标记弃用。autocast 只包前向和 loss,backward 放在上下文外。不要启用 autocast 后再全局 model.half,否则会绕开逐算子的安全策略。 算子策略不是“白名单里全降精度”。线性层、卷积通常进入 lower-precision;softmax、部分 loss 和归约倾向 FP32;多输入算子可能提升到最宽输入类型。显式 dtype、原地算子或 out 版本可能不参与 autocast,自定义代码要检查实际 dtype 流。 FP8 方面,NVIDIA Transformer Engine 提供 E4M3/E5M2 recipe、amax 历史与 delayed scaling,并在支持硬件上让 Transformer 层进入 FP8 路径。它仍以 BF16/FP16 保存部分张量并用高精度累加;分布式训练还可能跨 rank 归约 amax。“开启 FP8”是引入量化运行时,不是把参数永久存成单一 8 位格式。 生产监控至少记录当前 scale、增长和回退次数、跳过的 step 数、梯度范数、参数与激活 amax、NaN 首次出现的层、关键算子 dtype。只看总 loss 会把静默下溢藏很久。 06. 代价与边界 混合精度通常节省权重、梯度和激活带宽,但不保证总显存恰好减半。Adam 的 FP32 主权重与两个 moment 仍可能占 12 字节/参数;部分算子保存 FP32 中间量;通信 buffer、cast buffer、workspace 与碎片也增加峰值。 吞吐提升也有条件。矩阵尺寸太小、CPU 或数据加载受限、频繁转换、未使用低精度加速单元,都会让 AMP 几乎不加速。小 batch 下,kernel launch 与 cast 成本甚至抵消收益。正确基准应同时报告 tokens/s、峰值显存、收敛曲线和最终指标。 以下情况尤其要谨慎: 模型源自 BF16 预训练且激活常超过 65504,FP16 loss scaling 只能处理梯度,救不了前向溢出; 长归约、概率、指数、对数、方差或很小正则项的自定义算子没有 autocast 策略; 梯度累积时各 microbatch 使用不同 scale,或反缩放前裁剪梯度; 多 optimizer 共用计算图,却在所有梯度就绪前更新 scale; FP8 中异常值主导 amax,张量级缩放让普通值量化过粗,需要块缩放或高精度旁路。 调试时先用 FP32 建立可复现基线,再开 BF16,然后才是 FP16/FP8。每步只改一个变量,并比较前若干 step 的 loss、梯度范数和权重更新。数值问题最怕同时调整学习率、并行度、batch 与 dtype。 6.1 一次更新里究竟有哪些精度 以 AdamW 为例,一次更新至少涉及前向激活、反向梯度、参数、副本和两个动量。低精度 GEMM 只覆盖其中一部分。设低精度参数为 $\theta_{16}$,FP32 主参数为 $\theta_{32}$,反缩放后的梯度为 $g_{32}$,则更新链条可以写成: $$\theta_{32}^{t+1}=\operatorname{AdamW}(\theta_{32}^{t},g_{32}^{t},m_t,v_t),\qquad \theta_{16}^{t+1}=Q_{16}(\theta_{32}^{t+1})$$ $m_t,v_t$ 是 FP32 一阶、二阶矩,$Q_{16}$ 表示舍入到前后向 dtype。若框架没有主参数副本,而直接更新 BF16 参数,许多小于当前参数 ULP 的更新会消失。是否保留主权重是优化器实现细节,不能仅从模型参数 dtype 推断。 梯度累积又多一层顺序。假设一个 optimizer step 包含 $A$ 个 microbatch,它们必须使用同一 scale $S$,先把已缩放梯度累加: $$g_s=\sum_{a=1}^{A}Sg_a=S\sum_{a=1}^{A}g_a$$ 等所有 microbatch 完成后只反缩放一次。若中途 update scale,累加 buffer 中不同 microbatch 带有不同倍数,最后无法用一个 $S$ 恢复。若某个 microbatch 产生 inf,这一整个 optimizer step 都应跳过,而不是只丢掉坏 microbatch,否则有效 batch 和采样权重已变化。 6.2 分布式训练里的非有限值共识 数据并行中,每个 rank 只看到本地样本。rank 3 出现 inf 而其他 rank 有限时,所有副本仍必须对“是否更新”达成一致;否则 rank 3 跳步、其余 rank 更新,参数立即分叉。工业 GradScaler 会把 found-inf 状态跨相关设备和优化器汇总,分布式封装还需确保所有 DP rank 做相同决策。 梯度裁剪也有全局语义。ZeRO/FSDP 把梯度分片后,每卡只能算局部平方和: $$\|g\|_2=\sqrt{\sum_{r=0}^{P-1}\sum_{i\in\mathcal{S}_r}g_i^2}$$ 必须 all-reduce 局部平方和才能得到全局范数。正确顺序通常是:完成梯度同步,反缩放,检查非有限值,计算全局 norm,裁剪,执行 optimizer step,最后更新 scale。任一步调换都可能让“裁剪阈值 1.0”失去原含义。 混合精度还会改变 collective 的数值误差。BF16 梯度 all-reduce 比 FP32 少一半字节,但不同归约树改变加法顺序;卡数越多、梯度尺度跨度越大,末位差异越明显。若训练对归约误差敏感,可使用 FP32 梯度通信或分块高精度累加,但要接受带宽代价。排查时应分开比较“本地梯度生成精度”与“跨 rank 归约精度”。 6.3 一张可执行的 dtype 选型表 选 FP16 还是 BF16,不应只看硬件宣传峰值。可按四步判断: 第一,确认硬件原生路径。某些设备对 BF16 与 FP16 吞吐相同,某些旧设备没有 BF16 Tensor Core;软件模拟 BF16 可能更慢。第二,检查模型来源。BF16 预训练 checkpoint 的激活范围未必适合 FP16,直接转换容易前向溢出。第三,检查数值敏感区域。softmax、归一化、概率 loss、长归约优先保留 FP32。第四,跑短程 A/B,比较吞吐、峰值、跳步率和收敛,而非只确认“能启动”。 可以把训练方案分成四档: FP32 基线:最慢但诊断最清晰,用于建立 loss 与梯度参照; BF16 AMP:大范围、无需常规 loss scaling,是现代大模型训练的常见起点; FP16 AMP:尾数稍细,但指数窄,必须认真监控动态 scaling 与前向溢出; FP8 混合精度:只把适合的 GEMM 张量降到 FP8,保留 16/32 位旁路,并维护尺度元数据。 所谓“纯 BF16”或“纯 FP8”常是营销简写。softmax、norm statistic、optimizer state、某些 residual 累加和 master gradient 仍可能更高精度。真正有意义的配置描述,应列出 input、weight、output、accumulator 和 optimizer 五个维度。 6.4 从第一个 NaN 反推根因 排查顺序应尽量靠近异常源。先在每个模块前后记录有限值比例、绝对最大值、绝对非零最小值与 dtype;找到第一个从有限变非有限的边界。若 attention score 在 softmax 前已经 inf,检查 Q/K 范数、缩放因子与位置编码;若 loss 有限而 backward 首先 inf,检查导数奇点、自定义 backward 和 loss scale;若梯度有限但 step 后参数异常,检查 optimizer state、权重衰减和反缩放顺序。 静默下溢比 NaN 更难发现。建议记录零梯度比例,并按参数组看更新比率: $$r_{\mathrm{update}}=\frac{\|\Delta\theta\|_2}{\|\theta\|_2+\varepsilon}$$ 若某些层长期为 0,而 FP32 基线非零,说明低精度存储或 scale 窗口吞掉了更新。若所有层的比例突然增大,则更像学习率、loss normalization 或跳步恢复后的 scale 问题。 不要看到 NaN 就立即降低学习率。学习率过大确实可能导致发散,但 dtype overflow、错误 mask 产生全负无穷、空 batch 的除零、坏数据和通信错误都可能表现为同一个 NaN。先定位首次非有限张量,才能选择对应修复。 6.5 Checkpoint、编译与自定义算子的坑 恢复训练时必须保存 GradScaler 状态。若只恢复模型和 optimizer,却把 scale 重置为很大初值,前几步可能反复溢出并被跳过;重置太小则产生额外下溢。保存点最好位于完整 optimizer step 之后,避免记录到一半累积的混合状态。 activation checkpoint 会在 backward 重跑前向。重算必须进入与原前向一致的 autocast 上下文,否则保存路径是 BF16、重算路径变 FP32 或 FP16,梯度对不上。torch.compile、CUDA Graph 和 fused optimizer 还可能改变 autocast 边界或引入 CPU-GPU 同步,升级版本后要重新做数值与性能回归。 自定义 autograd Function 不能假定输入永远 FP32。前向若内部要求 FP32,应显式关闭 autocast 并转换输入;反向要返回与调用约定兼容的梯度。写自定义 CUDA kernel 时,明确 accumulator 类型、饱和或 inf 行为、次正规数处理以及随机舍入。一个算子“支持 half”只代表能执行,不代表适合训练。 最后,评估数值稳定不能只跑几十步。动态 scale 的增长周期可能是数千步,某些罕见 batch 才触发极端值。至少保留一次覆盖学习率峰值、warmup 结束和验证阶段的长程对照,并把跳步数作为训练产物记录;否则断点续训后出现指标差异,很难追溯是数据还是精度状态造成。 6.6 性能基准该怎样设计 AMP 基准至少要有三次 warmup 与多次稳定迭代,计时前后做设备同步;否则异步 kernel 会让 CPU 提交时间冒充 GPU 执行时间。显存要同时报告 allocated 与 reserved 峰值,前者是活跃张量,后者包含分配器缓存。对比方案必须使用同一 batch、相同梯度累积与相同 checkpoint 策略。 吞吐最好报告有效 token/s,而不是 batch/s。变长序列下,一个 batch 的 padding 比例不同,batch/s 会误导。质量侧至少比较训练 loss、验证指标、梯度范数和被跳过 step 数。若 AMP 每秒更快却需要更多 step 才到相同验证指标,最终 time-to-quality 未必更优。 对 FP8 还应增加量化覆盖率:多少 GEMM 真正走 FP8,多少因 shape、算子或 recipe 回退到 BF16;记录每层 amax、饱和比例和 scale 更新。仅看配置显示 FP8 无法证明加速路径被命中。kernel trace 能确认实际指令与 cast 开销。 建立门禁时,不必要求低精度与 FP32 每步完全相同。可先规定前 100 步 loss 相对误差、梯度 cosine、最终指标容差与最大跳步率,再用多个随机种子评估。数值差异是浮点并行的正常现象,持续偏向、层级爆炸或指标显著退化才是故障信号。 6.7 一份上线前检查清单 上线前逐项回答:硬件是否原生支持目标 dtype;矩阵尺寸是否对齐加速 tile;敏感算子是否保持 FP32;FP16 是否启用动态 scaling;裁剪是否在 unscale 后;所有 DP rank 是否共享跳步决策;梯度累积期间 scale 是否不变;checkpoint 是否保存 scaler;自定义算子是否声明 autocast 与 accumulator;监控是否能定位第一处非有限值。 然后做三个故障注入。人为把 scale 调得极大,确认系统发现 inf、跳过更新并回退;给输入加入一个幅值异常样本,确认前向监控能定位层;从 checkpoint 恢复,确认 scale、optimizer 与随机数状态连续。没有做过故障注入的告警,往往只在真正长跑失败后才发现无效。 最后保留 FP32 或 BF16 安全开关。生产训练发生异常时,能在不改数据顺序和并行布局的条件下提升精度复现,定位效率远高于临时改一堆超参。安全路径不一定长期运行,但必须定期测试,避免代码演进后早已失效。 这些流程看似比设置一个 autocast 开关繁琐,却能把“偶尔 NaN”“换卡就掉点”“断点后不收敛”变成有指标、有复现实验、有回退方案的普通工程问题。 6.8 如何判断变化来自精度而不是随机性 一次训练 A 比 B 的 loss 高,并不能证明精度方案更差。数据顺序、dropout、并行归约顺序和非确定 kernel 都会制造波动。公平实验要固定数据索引与初始化,尽可能使用确定算法,并运行多个随机种子。先比较同一步同一层的输出、梯度与更新,再比较长程最终指标。 可用 FP32 输出 $y_{32}$ 作为局部参照,计算相对误差与余弦相似度: $$e_{\mathrm{rel}}=\frac{\|y_{\mathrm{low}}-y_{32}\|_2}{\|y_{32}\|_2+\varepsilon}$$ 单层误差略大未必影响训练,但若误差沿深度单调放大,往往说明 residual 累加、归一化或某个敏感算子精度不足。把统计按层绘制比只比较最终 logits 更容易定位。 还应区分可复现性与正确性。集合通信改变加法顺序后,bitwise 结果可能不同,但两条训练曲线仍落在相同统计分布。反之,两次运行逐位一致也可能稳定地实现了错误缩放。验收既要有小规模数学 oracle,也要有多种子任务指标。 混合精度上线后,持续监控跳步比例和 scale 分布。数据配方变化、序列变长、加入新 loss 或更换初始化,都可能改变数值范围;一次验证通过不意味着未来配置永久安全。把 precision 当成模型配置的一部分进行版本化,才能复现每次训练。 一个实用原则是把所有隐式转换显式化到观测层:训练启动时抽样打印关键模块的参数、输入、输出与归约 dtype,同时保存硬件、驱动、框架版本。低精度 kernel 与 autocast 策略会随版本更新,旧实验结论不能无条件外推。版本升级后的第一件事应是重跑短程数值基线与吞吐基线,而不是直接续跑昂贵训练。这样才能知道速度变化来自 kernel,精度变化来自策略,还是数据本身发生了变化。 还应把这些元数据写进 checkpoint 清单,使恢复任务能够拒绝不兼容的精度配置,而不是在数小时后以 NaN 形式暴露。可诊断性本身就是混合精度系统的一项能力。 07. 经典论文脉络 Mixed Precision Training(arXiv:1710.03740):提出低精度前后向、FP32 主权重与 loss scaling,奠定现代 AMP。 A Study of BFLOAT16 for Deep Learning Training(arXiv:1905.12322):说明 BF16 用 FP32 相同指数范围换取更短尾数,使训练更少依赖 loss scaling。 FP8 Formats for Deep Learning(arXiv:2209.05433):提出 E4M3/E5M2 两种互补编码,并在大模型训练中验证 8 位浮点路径。 FP8-LM(arXiv:2310.18313):处理 FP8 大语言模型训练的精度、通信与优化器环节,说明端到端 FP8 不只是替换 GEMM dtype。 共同主题是:硬件格式越窄,软件越需要知道张量的数值角色。格式提供可能性,尺度管理、累加方式和回退路径才决定能否稳定训练。 08. 常见误解 误解一:BF16 比 FP16 精度更高。 BF16 指数范围更大,更少 overflow;但尾数只有 7 位,1 附近分辨率更粗。应说通常更稳定,不是处处更精确。 误解二:loss scaling 能修复所有 NaN。 它主要防 FP16 小梯度下溢。前向激活溢出、除零、softmax 不稳定、学习率过大,都不会被它根治。 误解三:scale 越大越好,而且一定大于 1。 scale 太大会溢出。PyTorch 允许动态 scale 降到 1 以下;应观察 found-inf 和跳步频率。 误解四:用了 autocast 就无需关心 dtype。 自定义算子、显式 dtype、原地运算和策略表外 op 可能保持输入 dtype。关键归约与 loss 仍需审计。 误解五:AMP 会把训练显存直接砍半。 激活可能下降,FP32 优化器状态仍在。账本还包括梯度、通信 buffer、workspace 与碎片。 误解六:FP8 是 BF16 的无痛升级。 FP8 依赖尺度、amax 历史、格式选择和硬件 kernel;错误配置可能不 NaN,却悄悄降低收敛质量。 09. 动手验证 先运行文末脚本,修改 loss_scaling_sim.py 中 scale 为 1、128、1024、65536。预期 scale 变大时更多微小梯度被保留;继续增大并加入大梯度后会触发 FP16 inf。这正是动态 scale 在 underflow 与 overflow 之间寻找窗口的原因。 在支持 BF16/FP16 的 GPU 上做四组短跑:FP32、BF16 autocast、FP16 autocast 不带 scaler、FP16 autocast 带 scaler。固定随机种子和 batch,记录 200 步吞吐、峰值显存、loss、全局梯度范数、跳步数。预期 FP16 无 scaler 最容易出现零梯度;BF16 与 FP16 加 scaler 更接近 FP32,但结论取决于模型分布。 定位首个异常层时,给模块注册 hook,输出有限值比例、绝对最大值和 dtype。若前向先出现 inf,优先检查指数、归一化、attention score 与输入范围;若反向先异常,再检查 scale、梯度裁剪顺序和自定义 backward。 最后验证累加精度:构造一个大数和大量小数的向量,用 FP32、模拟 BF16 顺序累加、分块 FP32 累加分别求和。预期低精度顺序加法吞掉更多小项;这解释为什么 GEMM 输入低精度不等于 accumulator 也该低精度。 10. 延伸阅读 读完这篇可以继续看: 量化:从 INT8 到 FP4:区分训练混合精度与推理量化的目标、校准和误差模型(还没写)。 数据并行与 ZeRO 显存切分:混合精度改变每项状态字节数,ZeRO 决定状态在哪些 rank 上复制或分片(已发布)。 张量并行与流水线并行:低精度减少通信字节,但集合通信次数、拓扑与同步点仍由并行策略决定(本文系列下一篇)。 可靠的 AMP 心智模型不是“半精度开关”,而是一张数值预算表:每个张量需要多大范围、多少有效位,在何处累加、何时舍入、出现异常如何回退。把这张表画清楚,NaN 就从玄学变成可以定位的工程问题。 附录:完整代码 09 节用到的脚本全文如下(float_formats.py、loss_scaling_sim.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 float_formats.py #!/usr/bin/env python3 """Only the Python standard library is required.""" import math import struct def fp16(x: float) -> float: try: return struct.unpack("e", struct.pack("e", x))[0] except OverflowError: return math.copysign(math.inf, x) def bf16(x: float) -> float: """Round an IEEE FP32 value to BF16, ties-to-even, then widen for printing.""" bits = struct.unpack(">I", struct.pack(">f", x))[0] if (bits & 0x7F800000) != 0x7F800000: bits += 0x7FFF + ((bits >> 16) & 1) return struct.unpack(">f", struct.pack(">I", bits & 0xFFFF0000))[0] FORMATS = ( ("FP16", 5, 10, 2.0**-14, (2 - 2.0**-10) * 2.0**15, 2.0**-10), ("BF16", 8, 7, 2.0**-126, (2 - 2.0**-7) * 2.0**127, 2.0**-7), ("FP32", 8, 23, 2.0**-126, (2 - 2.0**-23) * 2.0**127, 2.0**-23), ) print("format exp frac min_normal max_finite epsilon_at_1") for name, exp, frac, low, high, eps in FORMATS: print(f"{name:5s} {exp:3d} {frac:4d} {low:.4e} {high:.4e} {eps:.4e}") values = (1.0001, 0.00001, 100000.0) print("\nvalue -> FP16 | BF16") for value in values: print(f"{value:g} -> {fp16(value):g} | {bf16(value):g}") loss_scaling_sim.py #!/usr/bin/env python3 """Show how FP16 loss scaling rescues tiny gradients; stdlib only.""" import math import struct def fp16(x: float) -> float: try: return struct.unpack("e", struct.pack("e", x))[0] except OverflowError: return math.copysign(math.inf, x) gradients = [1e-8, 3e-8, 1e-7, 1e-6, 1e-5] scale = 1024.0 direct = [fp16(g) for g in gradients] scaled_then_unscaled = [fp16(g * scale) / scale for g in gradients] print(f"shape=({len(gradients)},), loss_scale={scale:g}") print("gradient direct_fp16 scaled_fp16/unscaled") for g, raw, rescued in zip(gradients, direct, scaled_then_unscaled): print(f"{g:11.3e} {raw:11.3e} {rescued:11.3e}") print(f"nonzero: direct={sum(x != 0 for x in direct)}/{len(direct)}, scaled={sum(x != 0 for x in scaled_then_unscaled)}/{len(direct)}") make_figures.py """Regenerate this article's deterministic teaching figure (numpy + matplotlib).""" from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np OUT=Path(__file__).resolve().parents[1]/"figures" OUT.mkdir(exist_ok=True) plt.rcParams.update({"font.sans-serif":["PingFang SC","Arial Unicode MS","DejaVu Sans"],"axes.unicode_minus":False}) from float_formats import fp16,bf16 x=np.geomspace(1e-9,1e5,1500) y16=np.array([fp16(float(v)) for v in x]);yb=np.array([bf16(float(v)) for v in x]) fig,ax=plt.subplots(figsize=(8,4.8)) for name,y in [("FP16",y16),("BF16",yb)]: err=np.abs(y-x)/x;err[~np.isfinite(err)]=np.nan ax.loglog(x,np.maximum(err,1e-12),label=name,alpha=.75) ax.axvline(65504,color="grey",ls=":",label="FP16 max finite") ax.set(xlabel="Positive input value",ylabel="Relative rounding error",title="Quantization range and precision (nearest-even)",ylim=(1e-8,2)) ax.legend();ax.grid(alpha=.25);fig.tight_layout() fig.savefig(OUT/'precision_error.png',dpi=170) plt.close(fig) print(OUT/'precision_error.png') 更多 AIGC 论文解读,关注微信公众号「人工智能炼丹君」 每日更新 · 论文精选 · 深度解读 · 技术脉络 微信搜索 人工智能炼丹君 或扫描下方二维码关注
2026年09月05日
7 阅读
0 评论
0 点赞
2026-09-05
AIGC 基本功|数据并行与 ZeRO 显存切分-ZeRO
数据并行与 ZeRO 显存切分 所属方向:分布式训练 | 难度:进阶 | 前置知识:无 关键词:数据并行、DDP、ZeRO、优化器状态切分、FSDP、显存占用估算 01. 为什么需要它 先算一笔会让很多人意外的账:训练一个 75 亿参数模型,模型权重明明只有 15 GB,为什么放进 80 GB 显存的 GPU 还会爆? 假设用混合精度 Adam 训练。每个参数除了 2 字节的 FP16/BF16 权重,还要留下 2 字节梯度、4 字节 FP32 主权重、4 字节一阶动量和 4 字节二阶动量。合计不是 2 字节,而是 16 字节/参数: $$7.5\times10^9\times16\ \mathrm{bytes}=120\ \mathrm{GB}$$ 这 120 GB 还没有算激活、临时通信缓冲区、CUDA context 和显存碎片。换句话说,模型甚至没开始处理一帧视频,单是“为了训练而保存的状态”就已经放不下。 自然的反应是:“那我上 64 张卡。”可如果只用普通数据并行,每张卡都会保留完整的参数、完整的梯度和完整的 Adam 状态。64 张卡只是把不同 mini-batch 同时算得更快,每张卡看到的仍然是同一份 120 GB。总显存从 80 GB 变成 5120 GB,但单卡瓶颈一点没松,模型照样启动不了。 这就是 ZeRO 要解决的矛盾:数据并行已经拥有整个集群的总显存,却因为每张卡都复制同样的训练状态,只能使用其中一张卡的容量。ZeRO 不改变模型的数学计算,而是按顺序切掉三种冗余副本: ZeRO-1 切优化器状态; ZeRO-2 再切梯度; ZeRO-3 连参数也切,只在某一层即将计算时临时拼回来。 把刚才的 75 亿参数模型放到 64 个数据并行 rank 上,三阶段的模型状态显存分别是 31.41 GB、16.64 GB 和 1.88 GB。最后一个数字正好是 $120/64$。这些数字不是宣传页上的模糊“最高节省多少”,而是本文后面会逐项推出来、再用代码复现的结果。 对视频生成尤其要补一句:视频 DiT 的 token 数很长,激活往往比模型状态还大。ZeRO 能拆掉模型状态冗余,但不会自动消灭激活。先分清是哪一类显存爆了,再决定用 ZeRO、激活重计算还是序列并行,比盲目把配置改成 stage 3 更重要。 02. 最小可用理解 三句话先建立框架: 数据并行让每个 rank 保存同一个模型、读取不同数据,反向后把局部梯度求平均,因而数学上等价于在合并后的大 batch 上做一次同步更新。 既然所有 rank 最终拿到相同梯度,它们各自保存一整套 Adam 状态并重复做同一份更新就是冗余;ZeRO 把状态分片,每个 rank 只负责其中 $1/N$。 切得越彻底,常驻显存越少,但参数就越需要“用前聚合、用后释放”,因此省下的显存会转化成通信、调度和峰值控制问题。 如果只记一件事:DDP 是复制计算,ZeRO 是切分状态;ZeRO 没有改变梯度本身。 这里的 rank 可以先理解为一个独立训练进程,通常一个 rank 独占一张 GPU。$N$ 表示数据并行进程数,world_size 就是 $N$。每个 rank 处理本地 batch,但每一步结束后所有 rank 必须得到一致参数,否则下一步就不再是同一个模型。 03. 数学推导 3.1 为什么平均梯度等价于一个大 batch 设全局 batch 有 $B$ 个样本,均匀切给 $N$ 个 rank,每个 rank 得到 $b=B/N$ 个样本。模型参数记作 $w$,第 $i$ 个样本的损失是 $\ell(x_i;w)$。 第 $r$ 个 rank 的局部平均损失为: $$L_r(w)=\frac{1}{b}\sum_{i\in\mathcal{B}_r}\ell(x_i;w)$$ 其中 $\mathcal{B}_r$ 是 rank $r$ 拿到的样本集合,$b$ 是本地 batch size。对参数求导,得到局部梯度: $$g_r=\nabla_w L_r(w)=\frac{1}{b}\sum_{i\in\mathcal{B}_r}\nabla_w\ell(x_i;w)$$ 同步数据并行对所有局部梯度取平均: $$g=\frac{1}{N}\sum_{r=1}^{N}g_r=\frac{1}{Nb}\sum_{r=1}^{N}\sum_{i\in\mathcal{B}_r}\nabla_w\ell(x_i;w)=\frac{1}{B}\sum_{i=1}^{B}\nabla_w\ell(x_i;w)$$ 右边正是把所有样本放在一张卡上算全局平均损失所得到的梯度。等价成立依赖三个前提:损失能按样本分解且样本计算不因分组改变(如本地 BatchNorm 统计会破坏这一条件);各 rank 从同一个参数 $w$ 出发;样本权重一致;更新前完成同步。若最后一个 batch 在各 rank 上不等长,仍然机械地“先本地平均再按 rank 平均”,样本就会被赋予错误权重。这也是分布式数据加载器必须谨慎处理尾 batch 的原因。 DDP 的核心任务由此非常明确:它不替你切输入,而是在 autograd 产生梯度后,把各 rank 的 $g_r$ 聚合成相同的 $g$。 3.2 先把 16 字节拆明白 设模型有 $P$ 个可训练参数,采用低精度前后向和 FP32 Adam 更新。每个参数对应: 低精度参数:2 字节; 低精度梯度:2 字节; FP32 主参数:4 字节; Adam 一阶矩 $m$:4 字节; Adam 二阶矩 $v$:4 字节。 后面三项统称优化器状态,共 12 字节。普通 DDP 每个 rank 的模型状态显存是: $$M_{\mathrm{DDP}}=(2+2+12)P=16P$$ 注意这是一个状态下界模型,没有包含激活等剩余显存。它的价值不是预测 nvidia-smi 到个位数,而是先回答“哪一类状态值得切”。 实际套公式时,$P$ 应取需要优化器更新的参数量,不是配置文件里笼统写出的模型总参数量。冻结的视觉编码器通常仍要保存推理权重,却不产生梯度和 Adam 状态;共享 embedding 在参数统计中也只能算一次。优化器同样会改常数:SGD 无动量时几乎没有 $m,v$,8-bit Adam 会压缩这两份状态,某些纯 BF16 配置也不保留 FP32 主参数。所以正确动作不是背住“训练恒等于 16 字节/参数”,而是先列出当前训练栈真正持有的张量,再把每一项的元素数乘 dtype 字节数。本文使用 16 字节,是为了与 ZeRO 原论文的混合精度 Adam 假设严格对齐。 3.3 ZeRO 三阶段到底切了什么 ZeRO-1:只切优化器状态。 参数和梯度仍然每卡各有一份,12 字节的 FP32 主参数与 Adam 两个矩按 $N$ 份切开: $$M_{\mathrm{Z1}}=4P+\frac{12P}{N}$$ 当 $N$ 很大,第二项趋近于 0,但第一项中的 2 字节参数和 2 字节梯度仍然复制,所以 ZeRO-1 的极限是 $4P$,相对 $16P$ 最多节省 4 倍。 ZeRO-2:再切梯度。 常驻的完整副本只剩 2 字节低精度参数;梯度与优化器状态合计 14 字节一起分片: $$M_{\mathrm{Z2}}=2P+\frac{14P}{N}$$ 当 $N$ 很大,极限是 $2P$,相对普通 DDP 最多节省 8 倍。这里“切梯度”不是说某个 rank 永远见不到别人的梯度,而是 reduce-scatter 聚合后,每个 rank 只保留自己负责更新的那一片结果,其余梯度用完即释放。 ZeRO-3:参数也切。 参数、梯度和优化器状态全部只保留 $1/N$: $$M_{\mathrm{Z3}}=\frac{16P}{N}$$ 这时单卡模型状态随卡数线性下降,理论上终于可以使用整个数据并行组的聚合显存。代价是某层计算前必须 all-gather 出完整参数,计算后再 reshard。这里的“完整”通常只针对一个 FSDP unit 或若干层,而不是一次把全模型永远拼回显存;包装粒度如果选错,峰值仍然可能爆掉。 代入 $P=7.5\times10^9$、$N=64$: DDP:$16P=120.00$ GB; ZeRO-1:$4P+12P/64=31.41$ GB; ZeRO-2:$2P+14P/64=16.64$ GB; ZeRO-3:$16P/64=1.875$ GB。 这与 ZeRO 原论文表 1中的 120、31.4、16.6、1.88 GB 对上了。论文使用十进制 GB;如果监控工具显示 GiB,同一字节数会小约 7.4%,比较数字前先统一单位。 3.4 显存省了,通信为什么没有立刻爆炸 设梯度张量大小为 $G$ 字节。高带宽 ring all-reduce 通常拆成 reduce-scatter 和 all-gather。忽略延迟、拓扑和协议常数后,每个 rank 在两个阶段分别移动: $$V_{\mathrm{RS}}=\frac{N-1}{N}G,\qquad V_{\mathrm{AG}}=\frac{N-1}{N}G$$ 因此普通 DDP 的总通信量近似为: $$V_{\mathrm{DDP}}=2\frac{N-1}{N}G\approx2G$$ ZeRO-1/2 把原来的 all-reduce 改排成“梯度 reduce-scatter + 更新后参数 all-gather”。在本文参数与梯度均为两字节的假设下,两者大小同为 $G$,理想总字节量仍约为 $2G$。若用 FP32 梯度而参数为 BF16,两项大小不同,不能直接沿用该字节比。区别在于数据留在哪里、何时释放,而不是凭空少传了一半。 ZeRO-3 还要在前向和反向各聚合一次参数,再做梯度 reduce-scatter。按 ZeRO 论文用参数元素数 $P$ 计量,普通数据并行约移动 $2P$ 个元素,Stage 3 约移动 $3P$ 个元素,所以是基线的 1.5 倍。这个 1.5 倍假设参数/梯度通信 dtype 相同并采用所述聚合与释放调度,是理想带宽模型,不是训练时间必然乘 1.5:通信能否和计算重叠、跨没跨节点、消息是否太碎,都会改变真实结果。 3.5 别把模型状态公式当成整机显存公式 实际峰值更接近: $$M_{\mathrm{peak}}=M_{\mathrm{model\ states}}+M_{\mathrm{activations}}+M_{\mathrm{temporary}}+M_{\mathrm{runtime}}+M_{\mathrm{fragmentation}}$$ 第一项就是前面推导的参数、梯度和优化器状态。激活由 batch、序列长度、空间分辨率、层数和 checkpoint 策略决定;临时项包含 all-gather bucket、梯度 bucket 和算子工作区;runtime 包含 CUDA context、通信库等;碎片则取决于张量生命周期和分配器状态。 ZeRO 原论文给过一个很有提醒意义的例子:15 亿参数 GPT-2 的模型状态至少 24 GB;序列长度 1024、batch 32 时,激活约 60 GB,即使用激活重计算降到约 8 GB,临时 FP32 扁平缓冲还可能再占 6 GB。只看 $16P$ 判断“32 GB 正好能放下”会直接翻车。 模型状态随数据并行卡数的变化。固定 FP16 参数/梯度与 FP32 主权重及 Adam 状态,纵轴为十进制 GB;不含激活、临时聚合缓冲和运行时。 04. 代码实现 本节有两个完整脚本,均只依赖 Python 标准库。第一个复现论文显存表;第二个把 DDP all-reduce 与 ZeRO-2 的 reduce-scatter + all-gather 拆开,验证它们做出同一次 Adam 更新。完整代码在文末附录。 4.1 用公式生成显存账本 zero_memory_ledger.py 把每个阶段写成“每卡复制多少字节 + 分片多少字节”: STAGES = ( Stage("DDP", 16.0, 0.0, 1.0), Stage("ZeRO-1", 4.0, 12.0, 1.0), Stage("ZeRO-2", 2.0, 14.0, 1.0), Stage("ZeRO-3", 0.0, 16.0, 1.5), ) def bytes_per_parameter(self, world_size: int) -> float: return self.replicated_bytes + self.sharded_bytes / world_size replicated_bytes 是每个 rank 必须完整保留的部分,sharded_bytes 是可以除以 $N$ 的部分。默认参数就是论文的 7.5B/64 卡案例。实际运行: model=7.5B, world_size=64 assumption: fp16 params 2B + fp16 grads 2B + fp32 master/m/v 12B = 16 bytes/parameter stage bytes/param model-state GB vs DDP comm/step GB DDP 16.0000 120.00 1.00x 30.00 ZeRO-1 4.1875 31.41 3.82x 30.00 ZeRO-2 2.2188 16.64 7.21x 30.00 ZeRO-3 0.2500 1.88 64.00x 45.00 note: activation, temporary buffers and fragmentation are excluded 为什么 64 卡的 ZeRO-1 只有 3.82 倍而不是宣传里的 4 倍?因为 4 倍是 $N\to\infty$ 的上限,有限卡数下 $12P/N$ 还没有消失。ZeRO-2 的 7.21 倍同理。ZeRO-3 没有复制项,所以恰好获得 64 倍。 再换成 1.5B/8 卡: model=1.5B, world_size=8 assumption: fp16 params 2B + fp16 grads 2B + fp32 master/m/v 12B = 16 bytes/parameter stage bytes/param model-state GB vs DDP comm/step GB DDP 16.0000 24.00 1.00x 6.00 ZeRO-1 5.5000 8.25 2.91x 6.00 ZeRO-2 3.7500 5.62 4.27x 6.00 ZeRO-3 2.0000 3.00 8.00x 9.00 note: activation, temporary buffers and fragmentation are excluded 这组数字揭示一个选型习惯:模型状态只差几 GB 时,ZeRO-1/2 往往已经够用;不必为了追求最低常驻显存直接承担 Stage 3 的参数聚合。 4.2 ZeRO 为什么没有改掉优化结果 zero_update_simulator.py 构造 4 个 rank、8 个参数。每个 rank 看见不同数据,因此产生不同局部梯度。DDP 路径先得到完整平均梯度,再完整执行 Adam;ZeRO-2 路径只把平均梯度对应的两元素分片交给各 rank,每个 rank 维护自己的 $m,v$,更新后再把四个参数片聚合起来。 核心区别只有分片发生的位置: reduced_full_grad = average_columns(local_grads) ddp_params, ddp_m, ddp_v = adam_first_step(PARAMS, reduced_full_grad) param_shards = shard(PARAMS, WORLD_SIZE) grad_shards = shard(reduced_full_grad, WORLD_SIZE) updated_shards = [] for params_for_rank, grads_for_rank in zip(param_shards, grad_shards): updated, local_m, local_v = adam_first_step(params_for_rank, grads_for_rank) updated_shards.append(updated) zero_params = [value for part in updated_shards for value in part] 真实运行输出: world_size=4, parameter_count=8, shard_width=2 local gradient matrix shape=(4, 8) reduce-scatter result shape=(4, 2) averaged gradient: [0.02, 0.045, 0.07, 0.095, 0.12, 0.145, 0.17, 0.195] DDP updated params: [0.099, -0.201, 0.299, -0.401, 0.499, -0.601, 0.699, -0.801] ZeRO-2 gathered params: [0.099, -0.201, 0.299, -0.401, 0.499, -0.601, 0.699, -0.801] max_abs_diff=0.000000000000 Adam m/v scalars per rank: DDP=16, ZeRO-2=4, reduction=4.00x max_abs_diff=0 是本文最重要的代码结果:同一份平均梯度、同一优化器规则下,把参数更新分给不同 rank 并不会改变更新后的模型。每卡 Adam 的 $m,v$ 元素数则从 16 降到 4,正好是 4 倍。真实系统还需要 bucket、异步通信、混合精度和异常恢复,但数学骨架就是这几十行。 05. 工业级实现对照 最小实现为了看懂“切什么”,生产实现要解决的则是“什么时候切、什么时候聚合、怎么让网络传输藏在计算后面”。下面以 2026-09-05 可见的官方实现为准。 5.1 PyTorch DDP:复制参数,反向时同步梯度 PyTorch 的 DistributedDataParallel 采用一进程一卡。它不会自动切分输入;应用通常用 DistributedSampler 保证不同 rank 读取不同样本。模型参数在每卡完整复制,autograd hook 在反向过程中把就绪梯度装进 bucket,并尽早发起 all-reduce,以便通信和后续层反向计算重叠。 这解释了两个常见现象:第一,DDP 通常比单进程 DataParallel 快,因为没有主卡收集和 Python 线程瓶颈;第二,DDP 加卡能缩短时间,却不会降低模型状态的单卡显存。它首先是一种吞吐扩展方案,不是大模型装载方案。 5.2 ZeroRedundancyOptimizer:最小改动的 Stage 1 思路 PyTorch 的 ZeroRedundancyOptimizer 可以和 DDP 组合。每个 rank 只为大约 $1/N$ 的参数维护本地优化器状态,更新自己负责的参数后广播结果,让所有 DDP 副本重新一致。这对应 ZeRO-1 的核心思想;参数仍完整复制,所以不能把它当成 FSDP 的替代品。 官方文档还标记该 API 为 experimental,并提示启用 overlap_with_ddp=True 时,最初若干迭代可能因为梯度 bucket 尚未稳定而不做参数更新。工程上不能只看“少了多少 GB”,还要读清 API 的更新时序、checkpoint 聚合和版本状态。 5.3 DeepSpeed ZeRO:三个阶段直接写进配置 DeepSpeed 官方 ZeRO 教程把三个阶段定义得很直接:Stage 1 切优化器状态,Stage 2 再切梯度,Stage 3 再切参数。一个典型配置是: { "zero_optimization": { "stage": 2, "contiguous_gradients": true, "overlap_comm": true, "reduce_bucket_size": 500000000 } } contiguous_gradients 针对碎片与通信连续性,overlap_comm 尝试把通信藏进计算,reduce_bucket_size 在“消息大到能吃满带宽”和“临时 bucket 不要撑爆显存”之间取舍。这些开关没有脱离前面的公式,只是在控制公式之外的临时项和时间轴。 如果 GPU 显存仍不够,ZeRO-Offload 可以把优化器状态与计算移到 CPU,ZeRO-Infinity 进一步使用 CPU 和 NVMe。但 官方 ZeRO-Offload 教程 展示的 10B 单 V100 案例并不意味着 offload 免费:PCIe 传输、CPU Adam 吞吐、NUMA 与磁盘带宽会变成新的瓶颈。 5.4 PyTorch FSDP:Stage 2/3 的框架化实现 知识树里的代码锚点 torch/distributed/fsdp/fully_sharded_data_parallel.py#FullyShardedDataParallel 在当前 PyTorch main 仍存在。FSDP1 的 ShardingStrategy 可以这样理解: NO_SHARD:参数、梯度、优化器状态都复制,行为接近 DDP; SHARD_GRAD_OP:梯度和优化器状态分片,参数在计算窗口内保持完整,接近 ZeRO-2; FULL_SHARD:参数、梯度和优化器状态全分片,接近 ZeRO-3; HYBRID_SHARD:节点内全分片、节点间复制,减少低带宽跨节点通信。 不过当前 PyTorch FSDP2 教程 已经明确建议迁移到 fully_shard:FSDP2 以 DTensor 做逐参数分片,不再依赖 FSDP1 的扁平参数;reshard_after_forward=True 对应 FULL_SHARD,设为 False 则更像 SHARD_GRAD_OP。文章或配置里只写“用了 FSDP”已经不够,必须同时说明代际和策略。 5.5 一层参数在 FSDP 中的生命周期 以 full shard 为例,一个 FSDP unit 在一次训练步里大致经历: 常驻状态只有本 rank 的参数分片; pre-forward hook 发起 all-gather,临时物化完整参数; 执行该 unit 的前向; 释放完整参数,恢复分片; 反向前再次 all-gather 参数; 计算梯度后执行 reduce-scatter,每个 rank 只留下本地梯度片; 本地优化器只更新本 rank 的参数与状态分片。 真正决定峰值的不是“最终只存 $1/N$”,而是同一时刻有多少 unit 正在 all-gather、预取队列有多深、最大 unit 有多大。如果把整个模型只包成一个巨型 unit,那么计算前仍要暂时物化全模型;如果把每个很小的算子都单独包起来,又会制造大量小消息,延迟和调度开销反而吞掉吞吐。Transformer 常按 block 包装,就是在这两端之间折中。 FSDP2 的公开契约仍是同一个时间逻辑:前向/反向前由 hook unshard,之后 reshard,梯度用 reduce-scatter 汇聚。实现细节从 flat parameter 变成 DTensor,不代表 ZeRO 的“按需物化”思想变了。 5.6 一个实用的选择顺序 先测量,再逐级加复杂度: 模型状态能放下,只想提吞吐:先用 DDP; Adam 状态是主要缺口:DDP + ZeRO-1/ZeroRedundancyOptimizer; 梯度也造成明显压力:ZeRO-2 或 FSDP 的 shard-grad 策略; 单卡连参数副本都放不下:ZeRO-3/FSDP full shard; 模型状态已经很小但长视频仍 OOM:处理激活重计算、micro-batch、FlashAttention 或序列并行,而不是继续折腾 ZeRO stage; 网络太慢:优先让分片组留在节点内,再用张量/流水线或混合分片扩到节点间。 06. 代价与边界 ZeRO 不减少总计算量。 同一个 batch 的前向、反向和 Adam 更新并没有少。它让每个 rank 少存状态,并通过通信在需要时恢复视图。若模型原本就能舒适放下,Stage 3 很可能只是增加通信和 hook 调度,吞吐反而下降。 通信字节相同,不代表时间相同。 ZeRO-2 与 DDP 在论文模型里都是约 $2G$,但一次大 all-reduce 和许多 layer-wise reduce-scatter/all-gather 的延迟特征不同。NVLink 节点内、InfiniBand 节点间、普通以太网的最优 bucket 大小不会一样;跨节点带宽不足时,1.5 倍 Stage 3 通信尤其明显。 平均显存很低,峰值仍可能 OOM。 参数 all-gather、预取、梯度归约、算子 workspace 可能同时在场。只看稳定阶段的 memory_allocated 会漏掉峰值;要记录 max_memory_allocated,并逐步调 wrap 粒度、prefetch 和 bucket。 ZeRO 解决不了激活随序列长度增长。 视频 DiT 中,时间、宽、高一起扩张,token 数可能成倍增加,注意力和 MLP 保存的激活随之增长。参数切到 1 GB 后仍然爆显存,并不说明 ZeRO 失效,而是瓶颈已经从模型状态转移到激活。 checkpoint 变复杂。 每个 rank 手里只有一片状态,保存时要选择 full、sharded 或 local state dict。把完整 checkpoint 聚合到 rank 0 可能让 CPU 内存瞬间成为瓶颈;只保存分片又要求恢复时正确处理 world size 和布局。训练能跑并不等于容灾链路可用,上线前至少做一次“保存—退出—换步数恢复”的演练。 全局操作必须认识分片。 梯度范数裁剪、参数检查、EMA、冻结部分参数、权重共享都不能默认“当前 rank 能看到完整张量”。例如 FSDP 提供自己的 clip_grad_norm_,就是因为全局范数需要跨分片归约。绕开框架直接遍历本地 .grad,算到的只是局部值。 offload 是拿带宽换容量。 CPU 内存比显存大,NVMe 又比 CPU 内存大,但层级越远,带宽越低、延迟越高。小模型或计算密度不够的模型会被数据搬运压垮;只有“不 offload 根本跑不了”或计算足以覆盖传输时,这个交换才划算。 大 batch 会改变优化问题。 增加数据并行度时,如果每卡 batch 不变,全局 batch 会随 $N$ 增长。ZeRO 保证同一全局 batch 下更新等价,却不保证扩大 batch 后收敛曲线不变。学习率、warmup、梯度累计和数据采样仍要一起调整。 什么时候不该用 ZeRO-3:模型与激活在单卡尚有充足余量、网络较慢、模型由大量极小模块组成且难以形成高效 bucket,或者你更看重最低延迟和调试简单性。此时 DDP 或 ZeRO-1/2 往往是更好的工程答案。 07. 经典论文脉络 这条路线不是“突然发明一种分布式训练”,而是一步步把数据并行里的冗余拆掉: Horovod(arXiv:1802.05799,2018)把 ring all-reduce 做成易接入的训练抽象,奠定了现代同步数据并行“各算各的、梯度集体归约”的工程基线;它解决吞吐扩展,但每个 worker 仍保存完整模型状态。 ZeRO(arXiv:1910.02054,2019)指出数据并行浪费的不是总显存,而是参数、梯度和优化器状态的重复副本;三阶段切分在保持数据并行计算粒度的同时,把模型状态显存从 $16P$ 推到 $16P/N$。 ZeRO-Offload(arXiv:2101.06840,2021)把优化器状态与计算搬到 CPU,并针对 CPU Adam 优化,让单张 32 GB V100 训练 10B 模型成为论文展示案例;留下的新问题是主机带宽和异构调度。 ZeRO-Infinity(arXiv:2104.07857,2021)继续把 CPU 与 NVMe 纳入内存层级,用带宽感知的分区和预取突破 GPU 内存墙;容量继续扩大,但 I/O 调度成为系统的核心。 PyTorch FSDP(arXiv:2304.11277,2023)总结了把 fully sharded data parallelism 纳入 PyTorch eager 训练栈的实践,让 ZeRO-3 类思想不再只属于单一外部训练引擎,并推动后续 FSDP2 的逐参数分片。 主线可以压缩成一句话:all-reduce 证明“计算可以复制、结果可以同步”,ZeRO 进一步问“既然结果会同步,状态为什么还要复制”。 后续 Offload、Infinity、FSDP 都是在回答状态放在哪、何时出现、以什么粒度搬运。 08. 常见误解 “数据并行会把模型切到多张卡上。” 不会。普通 DDP 是每卡一个完整模型,只切数据。能把吞吐从一张卡扩到八张,不代表能装下超过单卡容量的参数。 “ZeRO-3 后每卡永远只有 $1/N$ 参数。” 常驻状态是 $1/N$,计算某个 FSDP unit 前仍需临时 all-gather 完整参数。忽略这个瞬时窗口,就会得到理论显存能放、实际 forward 前仍 OOM 的配置。 “ZeRO-2 比 DDP 少传一半梯度,所以一定更快。” ZeRO-2 用 reduce-scatter 只留下梯度片,但更新后还要 all-gather 参数片。按论文的理想字节模型,总通信量与 DDP all-reduce 相同;快慢来自重叠、bucket、拓扑与实现,而非简单少一半。 “用了 ZeRO 就不用 activation checkpointing。” 两者切的是不同账本。ZeRO 处理模型状态冗余;activation checkpointing 用额外重算减少前向激活保存。长视频训练经常需要两者同时用。 “参数是 BF16,所以 Adam 也只占 2 字节。” 常见混合精度训练会保留 FP32 主参数与 FP32 的 $m,v$,优化器部分仍是 12 字节/参数。不同优化器、8-bit optimizer 或纯 BF16 更新会改变常数,因此使用 $16P$ 前必须先列出真实 dtype 与状态。 “64 卡 ZeRO-1 就一定省 4 倍。” 4 倍是 $N$ 很大时的渐近上限。64 卡的精确值是 $16/(4+12/64)=3.82$ 倍;8 卡只有 2.91 倍。宣传中的“up to”不能替代自己的显存账本。 “只要 loss 一样,所有数据并行更新都严格一致。” 浮点归约顺序会改变末位误差,随机数、dropout、数据尾 batch 和非确定性 kernel 也会影响复现。数学上等价不等于 bitwise identical;本文模拟得到 0 误差,是因为使用同一确定性顺序与双精度标量。 09. 动手验证 下面五个实验都能在没有 GPU 的机器上完成。两个脚本全文在附录,复制保存后直接用 Python 3 运行。 实验一:复现原论文表格。 python zero_memory_ledger.py 预期看到 7.5B/64 卡下 DDP、ZeRO-1/2/3 分别为 120.00、31.41、16.64、1.88 GB。如果结果差约 7%,检查自己是不是把 GiB 和十进制 GB 混用了。 实验二:观察有限卡数离理论上限有多远。 python zero_memory_ledger.py --params-b 1.5 --world-size 8 预期 ZeRO-1 只省 2.91 倍,ZeRO-2 只省 4.27 倍,ZeRO-3 才精确省 8 倍。再把 world-size 改成 2、4、16、64,观察 Stage 1 向 4 倍、Stage 2 向 8 倍收敛。 实验三:验证更新等价性。 python zero_update_simulator.py 预期 DDP 与 ZeRO-2 的八个新参数完全相同,max_abs_diff=0;每 rank 的 Adam $m,v$ 标量数从 16 降为 4。把 local_gradient() 的公式改掉,只要两条路径仍使用同一平均梯度,最终参数就应继续一致。 实验四:故意制造尾 batch 权重错误。 在模拟器里让最后一个 rank 只代表一个样本、其他 rank 各代表两个样本,然后比较“先对每个 rank 求平均再除以 4”和“按总样本数加权平均”。预期两个梯度不同。这能解释为什么真实训练要使用 DistributedSampler、drop_last 或正确的样本权重,而不能只相信 all-reduce。 实验五:给自己的模型补全账本。 从训练日志记录参数量、每种状态 dtype、峰值激活与临时 buffer。先用脚本算模型状态理论值,再和框架峰值相减。如果差值随序列长度或分辨率快速增长,瓶颈是激活;如果几乎不随输入变化,才优先继续切模型状态。预期这一步比直接试三个 ZeRO stage 更快找到真正的 OOM 原因。 10. 延伸阅读 读完这篇可以沿分布式训练知识树继续: 混合精度与数值稳定性:本文的 16 字节常数建立在混合精度 Adam 上;理解 FP16、BF16、FP8 的动态范围后,才能判断哪些状态可以继续降精度。 张量并行与流水线并行:当单个算子或单层连临时 all-gather 都放不下时,需要切计算本身;下一步要算 TP 通信量和 PP 气泡率。 序列并行与 Ring Attention:模型状态已经切完,长视频仍被激活卡住时,才轮到沿序列维度切注意力计算。 本文的核心资料均来自原始论文与官方实现:ZeRO 的显存和通信公式以 ZeRO 原论文为准;三阶段配置参考 DeepSpeed ZeRO 官方教程;DDP、优化器分片和 FSDP 策略分别核对了 PyTorch DDP 文档、ZeroRedundancyOptimizer 文档 与 FSDP 文档。代码路径和 API 会随上游重构,工业实现部分注明的日期就是版本边界。 附录:完整代码 09 节用到的脚本全文如下(zero_memory_ledger.py、zero_update_simulator.py、make_figures.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 zero_memory_ledger.py #!/usr/bin/env python3 """复现 ZeRO 论文中的混合精度 Adam 模型状态显存账本。 这里只计算参数、梯度和优化器状态,不包含激活、临时通信缓冲区、 CUDA context 与内存碎片。单位同时使用十进制 GB,便于和论文表格对照。 """ from __future__ import annotations import argparse from dataclasses import dataclass @dataclass(frozen=True) class Stage: name: str replicated_bytes: float sharded_bytes: float communication_multiple: float def bytes_per_parameter(self, world_size: int) -> float: return self.replicated_bytes + self.sharded_bytes / world_size STAGES = ( Stage("DDP", 16.0, 0.0, 1.0), Stage("ZeRO-1", 4.0, 12.0, 1.0), Stage("ZeRO-2", 2.0, 14.0, 1.0), Stage("ZeRO-3", 0.0, 16.0, 1.5), ) def gb(byte_count: float) -> float: """转为十进制 GB;ZeRO 原论文的 120 GB 使用这个口径。""" return byte_count / 1_000_000_000 def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("--params-b", type=float, default=7.5, help="参数量,单位十亿;默认复现论文的 7.5B 示例") parser.add_argument("--world-size", type=int, default=64, help="数据并行进程数") args = parser.parse_args() if args.params_b <= 0 or args.world_size <= 0: raise ValueError("params-b 与 world-size 必须为正数") params = args.params_b * 1_000_000_000 gradient_bytes = params * 2 # fp16/bf16 梯度 baseline_comm = 2 * gradient_bytes # ring all-reduce: reduce-scatter + all-gather print(f"model={args.params_b:g}B, world_size={args.world_size}") print("assumption: fp16 params 2B + fp16 grads 2B + " "fp32 master/m/v 12B = 16 bytes/parameter") print("stage bytes/param model-state GB vs DDP comm/step GB") ddp_bytes = STAGES[0].bytes_per_parameter(args.world_size) for stage in STAGES: bpp = stage.bytes_per_parameter(args.world_size) state_gb = gb(params * bpp) saving = ddp_bytes / bpp comm_gb = gb(baseline_comm * stage.communication_multiple) print(f"{stage.name:<8} {bpp:>10.4f} {state_gb:>17.2f} " f"{saving:>9.2f}x {comm_gb:>15.2f}") print("note: activation, temporary buffers and fragmentation are excluded") if __name__ == "__main__": main() zero_update_simulator.py #!/usr/bin/env python3 """用单进程模拟 DDP 与 ZeRO-2 的一次 Adam 更新。 脚本不依赖 GPU、PyTorch 或分布式运行时。四个“rank”先各自产生一份局部 梯度,再比较两条路径:DDP 对完整梯度做 all-reduce;ZeRO-2 对梯度做 reduce-scatter、每个 rank 只更新自己的参数片,最后 all-gather 参数。 """ from __future__ import annotations import math WORLD_SIZE = 4 PARAMS = [0.10, -0.20, 0.30, -0.40, 0.50, -0.60, 0.70, -0.80] LEARNING_RATE = 1e-3 BETA1 = 0.9 BETA2 = 0.999 EPSILON = 1e-8 def local_gradient(rank: int, width: int) -> list[float]: """构造确定性的局部梯度,模拟每个 rank 看见不同数据。""" return [((rank + 1) * (index + 2) - 3) / 100.0 for index in range(width)] def average_columns(rows: list[list[float]]) -> list[float]: return [sum(column) / len(rows) for column in zip(*rows)] def shard(vector: list[float], world_size: int) -> list[list[float]]: if len(vector) % world_size: raise ValueError("为了让示例清楚,参数量必须能被 world_size 整除") width = len(vector) // world_size return [vector[rank * width:(rank + 1) * width] for rank in range(world_size)] def adam_first_step(params: list[float], grads: list[float]) -> tuple[list[float], list[float], list[float]]: """执行 Adam 的第 1 步,并返回新参数、m、v。""" new_params: list[float] = [] m_state: list[float] = [] v_state: list[float] = [] for value, grad in zip(params, grads): m = (1.0 - BETA1) * grad v = (1.0 - BETA2) * grad * grad m_hat = m / (1.0 - BETA1) v_hat = v / (1.0 - BETA2) updated = value - LEARNING_RATE * m_hat / (math.sqrt(v_hat) + EPSILON) new_params.append(updated) m_state.append(m) v_state.append(v) return new_params, m_state, v_state def main() -> None: local_grads = [local_gradient(rank, len(PARAMS)) for rank in range(WORLD_SIZE)] # DDP:每个 rank 经 all-reduce 得到同一份完整平均梯度,并完整更新参数。 reduced_full_grad = average_columns(local_grads) ddp_params, ddp_m, ddp_v = adam_first_step(PARAMS, reduced_full_grad) # ZeRO-2:reduce-scatter 的结果等价于先平均,再只保留所属分片。 param_shards = shard(PARAMS, WORLD_SIZE) grad_shards = shard(reduced_full_grad, WORLD_SIZE) updated_shards: list[list[float]] = [] state_sizes: list[int] = [] for params_for_rank, grads_for_rank in zip(param_shards, grad_shards): updated, local_m, local_v = adam_first_step(params_for_rank, grads_for_rank) updated_shards.append(updated) state_sizes.append(len(local_m) + len(local_v)) # all-gather 后每个 rank 都能拿到同一份新参数;这里拼接一次代表该结果。 zero_params = [value for part in updated_shards for value in part] max_abs_diff = max(abs(a - b) for a, b in zip(ddp_params, zero_params)) shard_width = len(PARAMS) // WORLD_SIZE print(f"world_size={WORLD_SIZE}, parameter_count={len(PARAMS)}, " f"shard_width={shard_width}") print(f"local gradient matrix shape=({WORLD_SIZE}, {len(PARAMS)})") print(f"reduce-scatter result shape=({WORLD_SIZE}, {shard_width})") print("averaged gradient:", [round(x, 4) for x in reduced_full_grad]) print("DDP updated params:", [round(x, 6) for x in ddp_params]) print("ZeRO-2 gathered params:", [round(x, 6) for x in zero_params]) print(f"max_abs_diff={max_abs_diff:.12f}") print(f"Adam m/v scalars per rank: DDP={len(ddp_m) + len(ddp_v)}, " f"ZeRO-2={state_sizes[0]}, reduction={WORLD_SIZE:.2f}x") if __name__ == "__main__": main() make_figures.py """Regenerate this article's deterministic teaching figure (numpy + matplotlib).""" from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np OUT=Path(__file__).resolve().parents[1]/"figures" OUT.mkdir(exist_ok=True) plt.rcParams.update({"font.sans-serif":["PingFang SC","Arial Unicode MS","DejaVu Sans"],"axes.unicode_minus":False}) n=np.array([1,2,4,8,16,32,64]); p=7.5e9 fig,ax=plt.subplots(figsize=(8,4.8)) for name,y in [("DDP",np.full(n.shape,16.)),("ZeRO-1",4+12/n),("ZeRO-2",2+14/n),("ZeRO-3",16/n)]: ax.plot(n,y*p/1e9,"o-",label=name) ax.set(xscale="log",yscale="log",xlabel="Data-parallel ranks",ylabel="Model states per rank (GB)",title="7.5B model: FP16 weights/grads + FP32 master/m/v") ax.legend();ax.grid(alpha=.25);fig.tight_layout() fig.savefig(OUT/'zero_states.png',dpi=170) plt.close(fig) print(OUT/'zero_states.png') 更多 AIGC 论文解读,关注微信公众号「人工智能炼丹君」 每日更新 · 论文精选 · 深度解读 · 技术脉络 微信搜索 人工智能炼丹君 或扫描下方二维码关注
2026年09月05日
7 阅读
0 评论
0 点赞
2026-08-23
AIGC 基本功|视频生成中的强化学习与奖励模型-VideoRL
视频生成中的强化学习与奖励模型 所属方向:对齐与强化学习 | 难度:前沿专题 | 前置知识:DiffusionRL、FlowMatching 关键词:视频强化学习、奖励模型、Flow-GRPO、时序一致性奖励、人类偏好、reward hacking 01. 为什么需要它 先看一个用于说明奖励投机的假设场景;下面的“两千步”不是本文实际运行的视频实验记录。 你有一个文生视频模型,想用 RL 把它对齐到人类偏好。你手上有个打分器,在图像上验证过效果不错——清晰度、细节、构图都能打分。你把它接上视频,跑 GRPO,训了两千步。奖励曲线漂亮地往上走。 然后你去看生成结果:画面几乎不动了。 人物站在原地,衣角不飘,背景纹理纹丝不动。prompt 明明写着 "a person walking across the street",模型给你一张会呼吸的照片。 这不是 bug,是模型算得比你清楚。静止画面有一系列白拿的好处: 形状永远不会崩——因为它不变 帧间没有闪烁——因为帧之间完全一样 细节可以做到极精细——不用分算力去渲染运动模糊 而它放弃的只有「动态性」这一项。如果静止带来的加分超过动态性与指令遵循等全部损失,那么静止可能比运动拿到更高奖励,RL 会毫不犹豫地找到它。 视频奖励增加了运动与时序一致性维度。在有限模型能力下,运动幅度与形变、闪烁等伪影可能存在经验权衡;这不是“运动必然降低画质”的物理定律。图像的清晰度、构图、语义也可能冲突,不能假设图像奖励天然相容。 后面几节会把这笔账算成具体数字——用 VisionReward 真实发布的 29 个权重。 02. 最小可用理解 三句话: 采一组。对同一个 prompt 采样 $G$ 条视频(比如 8 条),用奖励模型给每条打一个分。 组内比。把这组分数减去组内均值、除以组内标准差,得到每条的优势(advantage)——比同组平均好就是正,差就是负。 按优势推。优势为正的往上推概率,为负的往下压。不需要训练价值网络(critic),组内均值就当基线用了。 这概括了 GRPO 的组内优势机制;完整目标还包含后文的概率比、裁剪和可选 KL。它省掉 critic 的代价是必须一次采一组,好处是少训一个网络、少一堆调不动的超参。 如果只记一件事:RL 不会给你「更好的视频」,只会给你「奖励更高的视频」。这两者的差距全部藏在奖励模型里。本文一半篇幅在讲奖励模型,不是跑题。 03. 数学推导 3.1 从策略梯度到组内基线 目标很直白:最大化生成样本的期望奖励。 $$J(\theta) = \mathbb{E}_{x \sim \pi_\theta} [r(x)]$$ 其中 $\pi_\theta$ 是生成模型(策略),$x$ 是一条采样出来的视频,$r(x)$ 是奖励模型给的分。对参数求梯度,用对数导数技巧: $$\nabla_\theta J = \mathbb{E}_{x \sim \pi_\theta} [r(x) \nabla_\theta \log \pi_\theta(x)]$$ 这个式子能用,但方差大得没法训。原因是 $r(x)$ 的绝对值直接乘在梯度上:假如所有样本的奖励都在 8.0 到 8.2 之间,那么每个样本都在说「往我这边推」,只是力度略有差别。真正有用的信息是「谁比谁好」,而不是「绝对分多高」。 标准解法是减一个基线 $b$: $$\nabla_\theta J = \mathbb{E}[(r(x) - b) \nabla_\theta \log \pi_\theta(x)]$$ 只要 $b$ 与 $x$ 无关,这个替换是无偏的(因为 $\mathbb{E}[b \nabla_\theta \log \pi_\theta] = b \nabla_\theta \mathbb{E}[1] = 0$)。PPO 那一路用一个学出来的价值网络当 $b$,GRPO 的选择更省事——同一个 prompt 采一组,用组内均值当基线: $$A_i = \frac{r_i - \mathrm{mean}(r_1, \dots, r_G)}{\mathrm{std}(r_1, \dots, r_G)+\varepsilon}$$ 这里不能直接套用上面的无偏证明:组均值含有 $r_i$,依赖当前样本。固定 prompt、独立同分布采样且不除标准差时, $$E[(r_i-\bar r)\nabla\log\pi_\theta(x_i)]=(1-1/G)\nabla J.$$ 使用不含自身的 leave-one-out 均值可消除该缩放;除以随机组标准差还会改变样本及 prompt 的相对权重。$\varepsilon>0$ 防止同分组除零,同分组的奖励优势为零。标准化是有用的工程选择,不能据此宣称完整 GRPO 无偏;奖励均值高本身也不等于期望梯度大。 这张图要看什么:左图是同一个 prompt 采出的 8 条视频的奖励,全挤在 8.0 到 8.2 之间——如果直接把 $r_i$ 乘在梯度上,8 条样本会一起喊「往我这边推」,有用的信息全被绝对值淹没(箭头方向完全一致)。右图减掉组内均值、除以组内标准差之后,同样的 8 个数字变成了有正有负的优势:绿的往上推、红的往下压。这就是 GRPO 不用 critic 也能训的原因——基线是从同组样本里白捡的(图由 code/make_figures.py 生成,下同)。 GRPO 出自 DeepSeekMath(arXiv:2402.03300),原本是给语言模型的数学推理用的,后来被搬到视觉生成上。 3.1.1 完整目标里还有两个护栏 上面那个式子是骨架。真实的 GRPO 目标还包着两层保护,后面 04 节的最小实现会把它们省掉,这里先说清省的是什么。 第一层是重要性采样比的截断(沿用 PPO 的做法): $$\mathcal{L} = -\mathbb{E}\left[ \min\left( \rho_i A_i, \ \mathrm{clip}(\rho_i, 1-\epsilon, 1+\epsilon) A_i \right) \right]$$ 式子里的 $\rho_i$ 是新旧策略的概率比,$\rho_i = \pi_\theta(x_i) / \pi_{\theta_{\text{old}}}(x_i)$。采样用的是旧策略 $\pi_{\theta_{\text{old}}}$,更新的是新策略 $\pi_\theta$,$\rho_i$ 修正这个偏差。clip 截断的是朝有利方向继续外推的收益,不保证每个 ratio 都在区间内,也不是 KL 的硬约束。神经网络共享参数,仍可能越界;需结合 KL 和梯度日志监控。视觉去噪实现通常按每个随机转移计算 ratio,不能把这里的整段视频记号误当成可直接算出的边缘密度。 第二层是对参考模型的 KL 惩罚: $$\mathcal{L}_{\text{total}} = \mathcal{L} + \beta \, \mathrm{KL}\left[ \pi_\theta \, \| \, \pi_{\text{ref}} \right]$$ $\pi_{\text{ref}}$ 通常是 RL 开始前的初始模型。这一项把策略拴在原始分布附近,是对付 reward hacking 最基础的护栏——因为大多数 hack(比如 01 节那个静止画面)都要求策略跑到离原始分布很远的地方去。$\beta$ 调的就是「允许它跑多远」。 代价是 $\beta$ 很难调:太大则学不动,太小则护栏形同虚设,而合适的值随任务和奖励模型变化。这也是为什么 04 节要先把护栏拆掉看清裸的优化行为——理解 hacking 怎么发生,比直接套护栏更重要。 3.2 flow matching:怎样得到随机转移的似然 Flow matching 的 ODE 为 $dx_t/dt=v_\theta(x_t,t)$。本文用 $t=1$ 表示噪声、$t=0$ 表示数据:给定初始噪声,逐步转移是确定性的条件 delta,不能直接代入高斯转移 log-prob 做 PPO/GRPO。 这不等于输出没有概率分布。随机初值经正则、可逆的 ODE 流仍可产生具有边缘密度的输出,CNF 可用散度积分计算密度;可微奖励还可以沿 ODE 用路径导数优化,确定性策略梯度也有自己的适用条件。受限的是本文这条“随机逐步转移似然 + score function”的路线。 Flow-GRPO 第 4.2 节 为它构造反向时间 SDE: $$dx_t=\left[v_\theta(x_t,t)-\frac{\sigma_t^2}{2}\nabla\log p_t(x_t)\right]dt+\sigma_t\,d\bar W_t,\qquad dt<0.$$ 由于时间从 1 向 0 走,score 修正前是减号;离散噪声标准差为 $\sigma_t\sqrt{-dt}$。若改用从噪声出发递增的时间坐标,漂移和 score 符号必须一起变换。准确 score、匹配的速度场与连续时间求解条件下,SDE 和 ODE 共享边缘分布;近似网络、有限步长及 RL 更新会带来偏差,不能保证实现后精确不变。 这使被训练的非退化转移具有可计算的高斯 log-prob。终步若噪声为零,仍要排除或另行处理,不能对零方差求对数密度。Flow-GRPO 还用 Denoising Reduction 减少训练采样步数、保留完整推理步数;这是论文所测任务上的经验结果,需在视频任务重新验证。 3.3 奖励从哪来:把打分变成一串是非题 现在缺的是 $r(x)$。视频的「好」是多维的,而人类偏好数据通常只有「A 比 B 好」这种成对比较,信息量很稀疏。 VisionReward(arXiv:2412.21059)的做法是把打分拆成一串可解释的是非题,再线性加权: $$r(x) = \frac{1}{K} \sum_{k=1}^{K} w_k \cdot a_k(x), \quad a_k(x) \in \{+1, -1\}$$ 其中 $a_k$ 是一个视觉语言模型对第 $k$ 个问题的回答(yes 记 $+1$,no 记 $-1$),$w_k$ 是该问题的权重。视频版一共 $K = 29$ 个问题,从「是否满足 prompt 的全部要求」到「细节是否精细」。 这个设计有两个好处。可解释:分低的时候能看出是哪一维拖的,不是一个黑盒数字。可控:想让模型更重视某一维,直接改 $w_k$ 就行。 代价是它把「好」压成了一个线性组合。而线性组合最容易被套利——下一节就用它真实发布的权重,把套利空间算出来。 04. 代码实现 4.1 先把奖励函数照抄一遍 VisionReward 的打分逻辑短到可以完整贴出来。这是它仓库里 inference-video.py 的核心两行(2026-08 的 main 分支): answers = np.array([1 if answer == 'yes' else -1 for answer in answers]) return np.mean(answers * weight).item() 有个细节值得停一下:no 记的是 $-1$,不是 $0$。这让答错变成实打实的扣分,而不只是「没拿到加分」。同一个维度上,yes 和 no 之间差了 $2 w_k$ 而不是 $w_k$——算套利空间时这个因子 2 不能漏。 29 个权重是仓库里 VisionReward_Video/weight.json 直接给出的。先看看它们长什么样(code/reward_hacking.py): 29 个维度,权重总和 6.2772,均值 0.2165 最高 1.1418(指令-未完全失败) 最低 0.0085(细节-不粗糙) 最高/最低 = 134 倍 指令遵循三档合计 2.3486,占总权重 37.4% 构图+色彩合计 0.0924,占总权重 1.5% 两个数字值得盯着看。最高和最低差 134 倍——这不是一个「各维度都重要」的均衡奖励,而是有着强烈倾向的。指令遵循三档占了 37.4%,构图加色彩只占 1.5%:在真实的人类偏好数据里,指令三档的权重合计是构图与色彩权重合计的约 25 倍。但这不是人类把指令遵循看得比整体美感重 25 倍的因果结论:其他维度也反映画质,且题目相关性、回答频率与权重尺度会影响解释。 这个结构反直觉但合理——观众要的是他要的东西,不是一张漂亮但无关的图。 这张图要看什么:左图把 29 个权重从低到高排开——最高的「指令-未完全失败」1.1418,最低的「细节-不粗糙」0.0085,相差 134 倍(红色是指令遵循三档,灰色是构图加色彩)。右图按语义归堆后更刺眼:指令遵循合计 2.3486、一家吃掉 37.4%,构图加色彩只有 0.0924、占 1.5%。指令三档的权重合计是构图与色彩这两组的二十多倍——记住这个比例,4.2 和 4.4 节的套利与翻转全是从它推出来的。 4.2 静态套利的账 回到 01 节那个静止画面。它能白拿哪些分、要放弃哪些分,现在可以精确算: 静止能白拿:全程形状保持+不混乱+画质稳定 = 0.7166 静止要放弃:镜头高度动态+不微弱 = 0.1634 净收益 +0.5532 (yes/no 相差 2 个单位,实际影响 ×2 = +1.1064) 上述 +1.1064 是未除以 $K$ 的加权和变化;score() 实际变化是 $1.1064/29\approx0.03815$。权重总和 6.2772 对应分数上限 0.2165,分数全范围宽度约 0.4329,因此这笔差额约占全范围 8.8%。这只是所选维度的局部账,不能省略指令等其他维度后宣称静止全局最优。 那它会不会因为违背 prompt 而被罚回来?我构造了两个视频画像对比(prompt 要求「人物走过街道」): 视频 A(几乎静止的精致画面,没照做 prompt) +0.1220 视频 B(照做了 prompt,但有形变和抖动) +0.1144 奖励模型更偏好:A(静态) 差值 0.0076 静态那个赢了。 拆开看差异来自哪: -1.9088 指令-全部满足 +0.8586 细节-非常精细 +0.5372 画质-非常稳定 +0.5272 全程形状-不混乱 -0.5048 指令-大部分满足 +0.3688 全程形状-完美保持 静态视频在指令遵循上亏了 2.41(这是权重最高的维度!),但靠画质那几项一点点堆回来,最后反超 0.0076。 需要说清楚:这两个 yes/no 画像是我按各维度语义人工假设的,不是真实模型的输出,所以这个结论说明的是「权重结构允许这样的套利」,不是「VisionReward 实际会这么判」。差值只有 0.0076 也说明它非常接近临界——是否会被优化器实际找到,还取决于策略可达性、采样噪声与正则约束,不能由这一对假设画像推断必然收敛。 4.3 让 GRPO 自己去找套利 手工构造画像终究不如让算法自己找。下面是完整的单参数 GRPO(code/grpo_minimal.py),策略只有一个参数:运动强度 $m$。 维度压到 1 维是刻意的——才能看清 GRPO 究竟把策略推向哪里。玩具概率按人为假设设定:$m$ 越大,指令遵循和动态性越容易拿 yes,形状保持、画质稳定、细节精细越容易拿 no。 def grpo(steps=400, group=32, lr=0.35, sigma=0.15, theta0=0.0, seed=0, weight=None): rng = np.random.default_rng(seed) theta = theta0 for step in range(steps): # 策略:m = clip(sigmoid(theta) + noise),高斯探索 mu = 1 / (1 + np.exp(-theta)) eps = rng.normal(0, sigma, group) ms = np.clip(mu + eps, 0.0, 1.0) rs = np.array([reward(m, rng, weight) for m in ms]) # 组内相对优势:减均值除标准差(3.1 节那个式子) adv = (rs - rs.mean()) / (rs.std() + 1e-8) # ∂mu/∂theta = mu(1-mu) grad = float(np.mean(adv * eps / sigma ** 2) * mu * (1 - mu)) theta += lr * grad return 1 / (1 + np.exp(-theta)) m_final = grpo() print(f"GRPO 收敛到 m = {m_final:.4f}") # GRPO 收敛到 m = 0.9657 十几行,没有 critic,也刻意去掉了 3.1.1 节那两层护栏(clip 和 KL)——保留的是组标准化 REINFORCE 示意,不是完整 PPO/GRPO 更新。代码中的高斯 score 对应裁剪前潜在动作 $a=\mu+\epsilon$,环境再执行 $m=\mathrm{clip}(a,0,1)$;不能把它当成裁剪后混合分布的普通高斯 log-prob。用完整的 29 个权重跑: step 运动强度 m 组内平均奖励 0 0.5000 -0.0238 80 0.9227 +0.0395 160 0.9420 +0.0346 320 0.9649 +0.0452 399 0.9660 +0.0370 GRPO 收敛到 m = 0.9657 → 策略选择:动态(照做了 prompt) 没有 hack。 策略从中性的 0.5 一路推到 0.966,选择了照做 prompt。换三个不同初值(0.27 / 0.50 / 0.73)出发,都收敛到 0.977 附近,说明这不是初值凑巧。 原因就是 4.1 节那个 37.4%:指令遵循的权重足够大,压住了画质维度的全部诱惑。这是玩具里的结果,不能验证 Flow-GRPO 的视频表现。Flow-GRPO 在其图像实验中他们报告「very little reward hacking occurred」,奖励涨上去了而画质和多样性没有明显退化。 4.4 关键对照:把指令遵循拿掉 那 01 节那个失败场景怎么发生的?答案在奖励模型上,不在算法上。 01 节的设定是「手上有个在图像上验证过的打分器」——图像打分器很可能根本不看时序,或者对指令遵循的权重远没有这么高。把三个指令遵循维度的权重置零,模拟这种情况,同一个 GRPO 再跑一遍: step 运动强度 m 组内平均奖励 0 0.5000 -0.0016 80 0.0305 +0.0269 160 0.0178 +0.0363 320 0.0099 +0.0396 399 0.0084 +0.0324 收敛到 m = 0.0084(完整权重下是 0.9657) 从 0.966 翻转到 0.008。 策略选择了几乎完全静止。注意中间那列奖励——它一路从 $-0.0016$ 涨到 $+0.04$,总体上也比初始值高,但逐步有明显随机波动。这就是 01 节说的:奖励曲线不能告诉你模型学到了什么。 奖励曲面看得更清楚——同一组物理约束,两种权重下的形状完全相反: m 完整权重 去掉指令遵循 0.0 -0.0393 +0.0393 0.3 -0.0298 +0.0228 0.5 +0.0006 -0.0006 0.8 +0.0446 -0.0322 1.0 +0.0379 -0.0394 → 完整权重的最高点在 m = 0.8(+0.0446) → 去掉指令遵循后最高点在 m = 0.0(+0.0393) 完整权重下曲面在 $m = 0.8$ 处见顶;去掉指令遵循后,在这一玩具设定下整体偏向低运动强度;表格是每点 4000 次随机采样的估计,不能仅凭稀疏网格确立精确单调性或全局最优。 这张图要看什么:左图是同一组物理约束下的奖励曲面——完整权重(蓝实线)在 $m=0.8$ 附近见顶,去掉指令遵循三档后(红虚线)整条曲线翻转成单调递减,最高点直接落到 $m=0$。右图是同一个 GRPO、同一个初值跑 400 步:蓝线冲到 0.966(照做 prompt),红线掉到 0.008(躺平)。而两条浅色细线是各自的组内平均奖励——二者总体奖励改善,但都有采样噪声,不能靠奖励曲线判断是否符合需求。奖励曲线不能告诉你模型学到了什么。 结论落在这里:reward hacking 不是算法出了毛病,而是它精确地最大化了你真正写下来的那个目标。 同一个 GRPO、同一个物理约束,只改奖励的权重结构,行为就翻转了。查 hacking 该去查奖励函数,而不是调 RL 的超参。 还要区分两种目标:图中的曲面是在固定运动强度 $m$ 上的采样均值,策略优化的却是高斯探索后再裁剪的期望奖励。$\mu=0.966$ 是裁剪前策略均值参数,不是实际运动强度的期望;它与网格中 $m=0.8$ 的峰值不能直接比较。因此这些数字不能证明“策略梯度过冲”,更不能归因为探索噪声必然造成方向性偏置。 05. 工业级实现对照 参考实现(以 2026-08 的 main 分支为准,上游会重构): THUDM/VisionReward → inference-video.py:score、compare_two_videos THUDM/VisionReward → VisionReward_Video/weight.json:29 个线性权重 上面的最小实现把奖励当成一个现成函数,生产上它是个 VLM,差别都在这。 5.1 打分要跑 29 次前向 score() 的实际流程是:抽帧、对每个问题拼一次 prompt、让 VLM 输出 yes/no。29 个问题就是 29 次 VLM 前向。 这直接决定了 RL 训练的成本结构。GRPO 每步要采一组(比如 8 条视频),每条视频要 29 次 VLM 前向,也就是每个训练步 232 次前向——而这还没算生成那 8 条视频本身的去噪开销。奖励计算不是附带开销,它经常是训练循环里最贵的一环。这也是 Flow-GRPO 要做 Denoising Reduction 的现实压力来源。 5.2 两套问题清单,别拿错 仓库里有两个文件长得很像: VisionReward_video_qa.txt:64 个问题,用于细粒度查询(问模型某一项如何) VisionReward_video_qa_select.txt:29 个问题,用于打分 inference-video.py 顶部读的是 _select 那个: QUESTIONS_PATH = "VisionReward_Video/VisionReward_video_qa_select.txt" WEIGHT_PATH = "VisionReward_Video/weight.json" 29 条问题对应 29 个权重,一一对齐。拿 64 个那份去乘权重会直接维度不匹配——这算好事,至少会报错。 5.3 比较两个视频不是比分数 想判断 A 和 B 哪个好,直觉是各算一次 score() 再比大小。它没这么做——compare_two_videos() 的结尾是: answers1 = np.array([1 if answer == 'yes' else -1 for answer in answers1]) answers2 = np.array([1 if answer == 'yes' else -1 for answer in answers2]) diff = answers1 - answers2 return np.sum(diff * weight).item() > 0 先求两个答案向量的逐维差 diff,再加权求和判正负。 只要两次打分使用同一组答案与权重,这与比较两个 score() 完全等价(仅差正的常数 $K$),并不能阻止不同维度相互抵消。VisionReward 论文的“多维一致性”另指偏好优化时筛选在各语义维度都一致占优的样本对;不是把减法移到加权和里面。 5.4 那些数字:这套设计到底有多少提升 VisionReward 报告的两个结果: 偏好预测准确率比 VideoScore 高 17.2% 用它做奖励的文生视频模型,pairwise win rate 比用 VideoScore 高 31.6% 这两个百分比对应不同指标与实验,不能相除推断“奖励改进被 RL 放大了”。VisionReward 的偏好优化还包括 DPO 等流程;本文引用的是论文报告结果,不是视频 GRPO 的对照实测。 06. 代价与边界 杠杆是双向的。 持续优化可能放大奖励偏差,但幅度不是由 5.4 节的两个异质指标决定的。奖励模型里一个不起眼的偏见——比如偏爱暖色调、偏爱中心构图——在评测里可能只是一两个百分点,但 RL 会把它当成目标,训几千步之后就是全部输出都发黄。奖励模型的偏见不会被平均掉,会被放大。 优化的是奖励,不是质量。 这句话看着像废话,但它有个实践推论:奖励曲线上升不能作为训练成功的证据。4.4 节那次 hack 的奖励曲线($-0.0016 \to +0.04$)和 4.3 节那次成功训练长得一模一样。 那怎么才能发现?三个手段,按性价比排: 盯住奖励模型看不见的指标。 挑几个不参与奖励计算的量化指标定期记录,比如帧间光流的平均幅度、逐帧 LPIPS 差异。静态 hack 在这类指标上一眼就露馅(光流幅度趋近 0),而奖励模型完全不看它们。这个办法便宜、能自动化,应该默认开着。 看奖励的维度分解,不只看总分。 VisionReward 这种线性结构的好处正在这里——把 29 维的得分分别记下来。总分涨、但某几维在持续下跌,就是套利正在发生。4.4 节第 5 段那张表就是这么读的。 固定一批 prompt 定期人眼看。 最贵但不可替代。前两个手段只能发现你预料到的失效模式,人眼能发现你没预料到的。 至于用第二个奖励模型交叉验证——听起来对称,实际有个陷阱:如果两个奖励模型是用同源的偏好数据训的,它们很可能共享同样的盲区,交叉验证会一起点头。要用就得选训练数据和架构都不同的。 成本很实在。 按 5.1 节的账,每个训练步是「$G$ 条视频的完整采样+ $29G$ 次 VLM 前向」。视频采样本身就比图像贵一个量级(多了时间维度),再乘上组大小,这是 RL 在视频上落地慢的主要原因。Denoising Reduction 这类工程手段不是优化项,是可行性前提。 线性加权是双刃剑。 它带来可解释性,也带来可套利性——线性函数没有交互项,意味着「这一维差到极点」不会拖累其他维度的加分。真实的人类判断不是这样的:一个完全静止的「视频」在人眼里直接不合格,不管它多精细。线性模型表达不了这种否决关系,所以 04 节那个套利在数学上成立。 什么时候不该用 RL。 如果你的问题能用更直接的手段解决,就别上 RL。想让输出更清晰——去修数据和 VAE;想让它听懂 prompt——先确认是不是文本编码器或者标注质量的问题。RL 适合的是「说得清好坏、但写不出损失函数」的目标,比如整体美感、运动自然度这类只能靠比较来表达的偏好。上 RL 之前先问一句:我的奖励模型真的比我的损失函数更懂这件事吗? 07. 经典论文脉络 这条线分两支:怎么优化和拿什么当奖励。两支交替推进。 奖励这一支: ImageReward(arXiv:2304.05977,2023-04)第一个规模化的文生图人类偏好模型,确立了「收集成对偏好数据训一个打分器」这个范式。它是个黑盒标量,可解释性问题从这里就留下了。 VideoScore(arXiv:2406.15252,2024-06)把细粒度人类反馈搬到视频上,成为视频奖励模型的常用基线——也就是 VisionReward 那两个数字(17.2% / 31.6%)超越的对象。 VisionReward(arXiv:2412.21059,2024-12)用分层是非题 + 线性加权换来可解释性,图像和视频统一处理。代价是线性结构本身可被套利(06 节)。 优化这一支: DDPO(arXiv:2305.13301,2023-05)把去噪过程当成多步决策,第一次让 policy gradient 在扩散模型上跑通。 DPOK(arXiv:2305.16381,2023-05)同期工作,补上了 KL 正则——防止策略为了刷奖励跑离原始分布太远。这是对付 hacking 最基本的护栏。 Diffusion-DPO(arXiv:2311.12908,2023-11)绕开 RL:直接用成对偏好数据做监督式优化,省掉在线采样。简单稳定,代价是只能利用已有的偏好数据,没有探索。 GRPO(arXiv:2402.03300,2024-02)出自 DeepSeekMath,本是给语言模型数学推理用的。用组内相对优势替掉 critic,把 RLHF 的工程复杂度砍掉一大块。 Flow-GRPO(arXiv:2505.05470,2025-05)第一个把在线 policy gradient 接进 flow matching 的工作,靠 ODE→SDE 转换取得随机逐步转移的可计算似然(3.2 节)。在文生图上把 GenEval 从 63% 拉到 95%、视觉文字渲染从 59% 到 92%。 需要说明:Flow-GRPO 的实验是在文生图上做的(SD3.5-M),不是视频。它之所以是视频 RL 的关键前置,是因为主流视频模型也是 flow matching 训的,ODE→SDE 那套推导可以直接搬过来——但视频上的实证还远没有图像充分。 08. 常见误解 「奖励涨了就是训好了」。 01 节整节都在讲这件事。奖励是你自己写的目标函数,模型最大化它是本分。奖励上升唯一证明的是优化在工作,不是效果在变好。 「reward hacking 是 RL 算法的问题,换个算法就好了」。04.4 节那个对照实验就是为了拆掉这个想法:同一个 GRPO,只改奖励权重,行为从 $m = 0.966$ 翻到 $m = 0.008$。算法只是执行者。出了 hacking 该去审奖励函数。 「GRPO 比 PPO 好」。 它主要是更省——不用训 critic,少一个网络和一堆超参。省掉 critic 的代价是必须一次采一组样本才能估出基线,采样开销更大。在视频这种采样极贵的场景里,这个权衡并不是无脑赚。 「多加几个奖励维度就能防 hacking」。 加维度确实提高了套利难度,但只要还是线性加权,套利空间就存在——4.2 节那个例子里有 29 个维度,静态依然套出了 1.1 的净收益。真正起作用的是权重结构(指令遵循占 37.4% 才压住了它),以及 DPOK 那种把策略拴在原始分布附近的 KL 护栏。 「flow matching 模型可以直接套用扩散模型的 RL 方法」。按随机转移似然做 PPO/GRPO 时,需要先定义非退化随机转移,例如 3.2 节的 ODE→SDE 路线。确定性 ODE 的输出边缘密度、可微奖励路径导数等另有方法,并不是所有 RL 都必须给每一步加噪。 09. 动手验证 三个实验,用到的脚本全文都在文末附录,纯 NumPy,没有 GPU 依赖,复制存成同名文件就能跑。 实验一:给自己的奖励函数算套利空间。 跑 reward_hacking.py,看第 2 段那笔账。然后把你自己项目里奖励各维度的权重填进 DIMS,重算一遍「静止能白拿 / 要放弃」的净收益。预期结论是:只要净收益为正,你的模型早晚会发现它。 实验二:找出翻转的临界点。 跑 grpo_minimal.py,把指令遵循的三个权重乘一个系数 $c$(改 weight_without_instruction 为按比例缩放),从 $c = 1.0$ 逐步降到 $0$,看收敛的 $m$ 在哪一步从 0.97 掉到 0.01。预期是个比较陡的转折而不是平滑过渡——这意味着奖励权重的微小改动可能导致行为质变,调权重时值得扫一遍而不是只试一个值。 实验三:验证组大小对方差的影响。 还是 grpo_minimal.py,把 group 从 32 改成 4、8、64,各跑三个不同 seed。预期看到组越小、收敛点越不稳定——因为组内均值和标准差是用这几个样本估的,样本太少基线就不可靠。还要一起权衡采样成本、同分组概率、奖励稀疏度和并行吞吐。 本文三张配图由 make_figures.py 生成(这个需要额外装 pip install matplotlib,其余脚本只要 NumPy)。想复现就把权重换成你自己的:左边的权重账本换成你的 29 维,右边的曲面和 GRPO 轨迹会直接告诉你你的奖励函数在纵容哪种躺平。 10. 延伸阅读 这篇在知识树里是 L4(前沿专题),前面还有三级台阶。如果 03 节读起来吃力,按下面顺序补更省时间: 策略梯度与 PPO 基础——3.1 节那个对数导数技巧、为什么要减基线、clip 是干什么的,都属于那一篇 从 DPO 到 GRPO:去掉价值网络——GRPO 本身的完整讨论在那篇。本文只用到它「组内相对优势」这一个机制,没有展开和 PPO、DPO 的对比 把 RL 用到扩散模型上——DDPO 那套「把去噪当多步决策」的建模 流匹配与 Rectified Flow——3.2 节 ODE→SDE 转换的前提 读完这篇可以继续看: 人类偏好数据怎么收(还没写)——4.1 节那 29 个权重是怎么从成对比较数据里拟合出来的 视频质量评测指标(还没写)——06 节说的那三个检测手段该怎么落实,尤其是「奖励模型看不见的指标」具体该记哪些 附录:完整代码 09 节用到的脚本全文如下(make_figures.py、reward_hacking.py、grpo_minimal.py)。复制到本地存成同名文件,按各脚本开头的依赖说明准备环境后即可运行。 make_figures.py #!/usr/bin/env python3 # -*- coding: utf-8 -*- """生成「视频生成中的强化学习与奖励模型」的三张解释图。 数值一律从同目录两个脚本里取(reward_hacking 的 29 个真实权重、 grpo_minimal 的奖励曲面与 GRPO 轨迹),这里只负责画——那边改了这里要重跑, 免得图和正文数字打架。 三张图分别回答: 1. 29 个奖励权重的账本长什么样(为什么「有没有照做」压过「好不好看」) 2. 组内优势到底把什么信息抽出来了(原始奖励看着都一样,标准化后才有正负) 3. 只改奖励权重、不动算法,GRPO 为什么会从「照做」翻转成「躺平」 只依赖 numpy + matplotlib。跑法:python make_figures.py """ import textwrap from pathlib import Path import matplotlib.pyplot as plt import numpy as np import grpo_minimal as G from reward_hacking import NAMES, WEIGHT 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_INSTR = "#d1495b" # 指令遵循:红 C_OTHER = "#2f6fb0" # 其余维度:蓝 C_AES = "#c9ced6" # 构图/色彩:灰 C_FULL = "#2f6fb0" # 完整权重 C_NOIN = "#d1495b" # 去掉指令遵循 C_POS = "#2f9e6f" # 优势为正 C_NEG = "#d1495b" # 优势为负 # 维度分组(下标对应 reward_hacking.DIMS 的顺序) GROUPS = [ ("指令遵循", [0, 1, 2], C_INSTR), ("构图+色彩", [3, 4, 6], C_AES), ("光照", [7, 8, 9, 10, 11, 12], C_OTHER), ("形状保持", [13, 14, 15, 16, 17], C_OTHER), ("运动", [5, 18, 19, 20, 21], C_OTHER), ("画质+细节", [22, 23, 24, 25], C_OTHER), ("文字+内容", [26, 27, 28], C_OTHER), ] 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:29 个奖励权重的账本 # ────────────────────────────────────────────────────────────────────── def fig_weight_ledger(): order = np.argsort(WEIGHT) # 从小到大,画出来是升序 colors = [C_INSTR if i in (0, 1, 2) else C_AES if i in (3, 4, 6) else C_OTHER for i in order] labels = [NAMES[i] for i in order] fig, axes = plt.subplots(1, 2, figsize=(12.6, 6.4), gridspec_kw={"width_ratios": [1.45, 1]}) fig.subplots_adjust(wspace=0.34) ax = axes[0] ax.barh(np.arange(29), WEIGHT[order], color=colors, height=0.72) ax.set_yticks(np.arange(29)) ax.set_yticklabels(labels, fontsize=7.6) ax.set_xlim(0, WEIGHT.max() * 1.18) ax.tick_params(axis="y", length=0) style(ax, "29 个维度的权重(升序)", "权重 $w_k$", None) # 最高 / 最低标注 ax.text(WEIGHT[order][-1] + 0.02, 28, f"{WEIGHT.max():.4f} 最高", fontsize=8.5, color=INK, va="center") ax.text(WEIGHT[order][0] + 0.02, 0, f"{WEIGHT.min():.4f} 最低", fontsize=8.5, color=MUTE, va="center") # 图例:红=指令遵循,灰=构图色彩,蓝=其余 from matplotlib.patches import Patch ax.legend(handles=[ Patch(color=C_INSTR, label="指令遵循(3 维)"), Patch(color=C_AES, label="构图+色彩(3 维)"), Patch(color=C_OTHER, label="其余(23 维)"), ], fontsize=8.2, frameon=False, loc="lower right") ax = axes[1] gcolor = {name: c for name, _, c in GROUPS} gs = [(name, float(WEIGHT[idx].sum())) for name, idx, _ in GROUPS] gs.sort(key=lambda t: -t[1]) total = WEIGHT.sum() ys = np.arange(len(gs)) # 颜色必须按排序后的名字取:按 GROUPS 顺序 zip 会跟排序后的柱子错位 ax.barh(ys, [v for _, v in gs], color=[gcolor[n] for n, _ in gs], height=0.7) ax.set_yticks(ys) ax.set_yticklabels([n for n, _ in gs], fontsize=9) ax.tick_params(axis="y", length=0) for i, (_, v) in enumerate(gs): ax.text(v + 0.03, i, f"{v:.3f} ({v / total * 100:.1f}%)", fontsize=8.2, color=INK, va="center") ax.set_xlim(0, max(v for _, v in gs) * 1.45) style(ax, "按语义归堆:谁在主导总分", "权重合计", None) fig.suptitle("图 1 VisionReward 的 29 个权重不是均衡的:" "指令遵循一家占了 37.4%,构图加色彩只有 1.5%", fontsize=12, color=INK, x=0.012, ha="left", y=1.02) footer(fig, "要看什么:左图是 29 个权重从低到高排开——最高的「指令-未完全失败」1.1418," "最低的「细节-不粗糙」0.0085,相差 134 倍,根本不是「各维度都重要」的均衡奖励。" "右图把它们按语义归堆后更刺眼:指令遵循三档合计 2.3486,一家吃掉 37.4%;" "而构图加色彩只有 0.0924、占 1.5%。真实人类偏好数据里," "「有没有照着 prompt 做」比「好不好看」重要二十多倍——这个结构是后面所有套利的源头。") fig.savefig(OUT / "reward_weight_ledger.png") plt.close(fig) print(f"[图1] 29 维权重合计 {total:.4f};最高 {WEIGHT.max():.4f}/最低 {WEIGHT.min():.4f}" f" = {WEIGHT.max() / WEIGHT.min():.0f} 倍;指令遵循占 " f"{WEIGHT[[0, 1, 2]].sum() / total * 100:.1f}%") # ────────────────────────────────────────────────────────────────────── # 图 2:组内优势把「谁比谁好」抽出来 # ────────────────────────────────────────────────────────────────────── def fig_group_advantage(): rng = np.random.default_rng(2024) rs = 8.0 + rng.random(8) * 0.2 # 同一个 prompt 采 8 条,分都挤在 8.0~8.2 adv = (rs - rs.mean()) / (rs.std() + 1e-8) idx = np.arange(1, 9) fig, axes = plt.subplots(1, 2, figsize=(12.6, 4.6)) fig.subplots_adjust(wspace=0.26) ax = axes[0] ax.bar(idx, rs, color=C_OTHER, width=0.62) ax.set_ylim(0, 8.9) ax.set_xticks(idx) style(ax, "原始奖励:8 条都在 8.0~8.2 之间", "组内第 i 条", "奖励 $r_i$") for i, v in zip(idx, rs): ax.arrow(i, v + 0.42, 0, 0.30, width=0.055, color=INK, length_includes_head=True, head_width=0.16, head_length=0.12) ax.text(0.99, 8.72, "每一条都在说:往我这边推(只是力度略有差别)", fontsize=8.6, color=MUTE, ha="right", va="top") ax.text(0.02, 0.30, f"极差只有 {rs.max() - rs.min():.2f}", fontsize=8.6, color=INK, transform=ax.transAxes) ax = axes[1] cols = [C_POS if a >= 0 else C_NEG for a in adv] ax.bar(idx, adv, color=cols, width=0.62) ax.axhline(0, color=INK, lw=1.0) ax.set_xticks(idx) ax.set_ylim(-2.0, 2.0) for i, v in zip(idx, adv): ax.text(i, v + (0.09 if v >= 0 else -0.09), f"{v:+.2f}", ha="center", va="bottom" if v >= 0 else "top", fontsize=8, color=INK) style(ax, "减均值除标准差后:谁该推、谁该压,一眼分清", "组内第 i 条", "优势 $A_i$") ax.text(0.98, 0.94, "绿的往上推,红的往下压;组内均值自动当基线", fontsize=8.6, color=MUTE, ha="right", va="top", transform=ax.transAxes) fig.suptitle("图 2 GRPO 的组内基线:把「绝对分多高」换成「谁比谁好」", fontsize=12, color=INK, x=0.012, ha="left", y=1.04) footer(fig, "要看什么:左图是同一个 prompt 采出的 8 条视频的奖励,全挤在 8.0 到 8.2 之间——" "如果直接把 $r_i$ 乘在梯度上,8 条样本会一起喊「往我这边推」,只是力气略有差别," "有用的信息被绝对值淹没了(箭头方向完全一致)。右图减掉组内均值、除以组内标准差之后," "同样的 8 个数字变成了有正有负的优势:一半往上推、一半往下压。" "这就是 GRPO 不用 critic 也能训的原因——基线是从同组样本里白捡的。") fig.savefig(OUT / "group_advantage.png") plt.close(fig) print(f"[图2] 8 条奖励 {rs.min():.3f}~{rs.max():.3f}(极差 {rs.max() - rs.min():.3f});" f"优势 {adv.min():+.2f}~{adv.max():+.2f}") # ────────────────────────────────────────────────────────────────────── # 图 3:只改奖励权重,同一个 GRPO 翻转 # ────────────────────────────────────────────────────────────────────── def _trace_grpo(weight=None, seed=0, steps=400, group=32, lr=0.35, sigma=0.15): """复刻 grpo_minimal.grpo 的循环,只是把每步的 m 和组内平均奖励记下来。 抽随机的顺序必须和原函数完全一致,否则数字对不上正文。 """ rng = np.random.default_rng(seed) theta = 0.0 ms, rr = [], [] for _ in range(steps): mu = 1 / (1 + np.exp(-theta)) eps = rng.normal(0, sigma, group) m_clip = np.clip(mu + eps, 0.0, 1.0) r = np.array([G.reward(m, rng, weight) for m in m_clip]) adv = (r - r.mean()) / (r.std() + 1e-8) grad = float(np.mean(adv * eps / sigma ** 2) * mu * (1 - mu)) theta += lr * grad ms.append(mu) rr.append(float(r.mean())) return np.array(ms), np.array(rr), 1 / (1 + np.exp(-theta)) def fig_surface_flip(): w_no = G.weight_without_instruction() # 奖励曲面:与 grpo_minimal.main 同口径(rng=7,每点 4000 次采样) rng = np.random.default_rng(7) grid = np.linspace(0, 1, 21) surf_full, surf_no = [], [] for m in grid: surf_full.append(float(np.mean([G.reward(m, rng) for _ in range(4000)]))) surf_no.append(float(np.mean([G.reward(m, rng, w_no) for _ in range(4000)]))) surf_full = np.array(surf_full) surf_no = np.array(surf_no) # 正文引用的那张 11 点表,峰值标注以它为准 rng2 = np.random.default_rng(7) tbl = np.linspace(0, 1, 11) tf, tn = [], [] for m in tbl: # 抽随机顺序与 grpo_minimal.main 完全一致(逐点交错),否则峰值数字对不上正文 tf.append(float(np.mean([G.reward(m, rng2) for _ in range(4000)]))) tn.append(float(np.mean([G.reward(m, rng2, w_no) for _ in range(4000)]))) tf, tn = np.array(tf), np.array(tn) pk_f, pk_n = tbl[int(tf.argmax())], tbl[int(tn.argmax())] ms_full, rr_full, m_end_full = _trace_grpo(weight=None, seed=0) ms_no, rr_no, m_end_no = _trace_grpo(weight=w_no, seed=0) fig, axes = plt.subplots(1, 2, figsize=(12.6, 5.0)) fig.subplots_adjust(wspace=0.30) ax = axes[0] ax.plot(grid, surf_full, lw=2.4, color=C_FULL, label="完整 29 维权重") ax.plot(grid, surf_no, lw=2.4, color=C_NOIN, ls="--", label="去掉指令遵循三档") ax.axhline(0, color="#9ca3af", lw=0.8) ax.scatter([pk_f], [tf.max()], s=70, color=C_FULL, zorder=5) ax.scatter([pk_n], [tn.max()], s=70, color=C_NOIN, zorder=5) ax.annotate(f"最高点 m={pk_f:.1f}", (pk_f, tf.max()), textcoords="offset points", xytext=(10, -30), fontsize=8.6, color=C_FULL, arrowprops=dict(arrowstyle="->", color=C_FULL, lw=1.0)) ax.annotate(f"最高点 m={pk_n:.1f}", (pk_n, tn.max()), textcoords="offset points", xytext=(18, 26), fontsize=8.6, color=C_NOIN, arrowprops=dict(arrowstyle="->", color=C_NOIN, lw=1.0)) ax.legend(fontsize=8.6, frameon=False, loc="lower center") style(ax, "同一组物理约束,两种权重下的奖励曲面形状完全相反", "运动强度 $m$", "期望奖励 $r$") ax = axes[1] ax.plot(ms_full, lw=2.4, color=C_FULL, label="m(完整权重)") ax.plot(ms_no, lw=2.4, color=C_NOIN, ls="--", label="m(去掉指令遵循)") ax.set_ylim(-0.05, 1.05) ax.set_xlabel("训练步", fontsize=10, color=MUTE) ax.set_ylabel("运动强度 $m$", fontsize=10, color=MUTE) ax2 = ax.twinx() ax2.plot(rr_full, lw=1.1, color=C_FULL, alpha=0.45) ax2.plot(rr_no, lw=1.1, color=C_NOIN, alpha=0.45, ls="--") ax2.set_ylabel("组内平均奖励(浅色细线)", fontsize=9, color=MUTE) ax2.tick_params(colors=MUTE, labelsize=8.5) ax2.spines["top"].set_visible(False) h1, l1 = ax.get_legend_handles_labels() ax.legend(h1, l1, fontsize=8.6, frameon=False, loc="center right") style(ax, "同一个 GRPO 跑 400 步:一条冲到 0.97,一条躺到 0.01", None, None) ax.set_ylabel("运动强度 $m$", fontsize=10, color=MUTE) ax.set_xlabel("训练步", fontsize=10, color=MUTE) ax.text(0.02, 0.06, f"完整权重收敛 m={m_end_full:.3f}\n" f"去掉指令遵循收敛 m={m_end_no:.3f}", fontsize=8.4, color=INK, transform=ax.transAxes, va="bottom", bbox=dict(fc="white", ec="#d1d5db", lw=0.7, pad=4)) fig.suptitle("图 3 奖励曲线一样漂亮,学到的东西完全相反", fontsize=12, color=INK, x=0.012, ha="left", y=1.03) footer(fig, "要看什么:左图是同一组物理约束下的奖励曲面。完整权重(蓝实线)在 m=0.8 附近见顶," "去掉指令遵循三档后(红虚线)整条曲线翻转成单调递减,最高点直接落到 m=0——彻底不动才是最优解。" "右图是同一个 GRPO、同一个初值跑 400 步:蓝线冲到 0.966 选择照做 prompt," "红线掉到 0.008 选择躺平。而两条浅色细线是各自的组内平均奖励——它们都在稳步上涨," "形状和一次成功的训练没有任何区别。奖励曲线不能告诉你模型学到了什么。") fig.savefig(OUT / "reward_surface_flip.png") plt.close(fig) print(f"[图3] 曲面峰值:完整 m={pk_f:.1f} ({tf.max():+.4f})," f"去掉指令 m={pk_n:.1f} ({tn.max():+.4f});" f"GRPO 收敛 {m_end_full:.4f} vs {m_end_no:.4f}") def main(): fig_weight_ledger() fig_group_advantage() fig_surface_flip() 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() reward_hacking.py """用 VisionReward 的真实权重,算一遍「静态视频能不能骗到高分」。 权重与问题取自 THUDM/VisionReward(2026-08 的 main 分支): VisionReward_Video/weight.json 29 个线性权重 VisionReward_Video/VisionReward_video_qa_select.txt 29 个 yes/no 问题 打分公式同 inference-video.py 的 score():yes=+1,no=-1,乘权重后取均值。 注意:下面两个视频的 yes/no 画像是**按各维度语义做的人工假设**, 不是真实模型的输出。它用来说明权重结构本身允许什么样的套利, 不构成对 VisionReward 实际表现的评测。 运行 `python reward_hacking.py`。 """ import numpy as np # (权重, 维度简称) —— 顺序与上游 weight.json 严格一致 DIMS = [ (0.9544, "指令-全部满足"), (0.2524, "指令-大部分满足"), (1.1418, "指令-未完全失败"), (0.0350, "构图-美观"), (0.0252, "构图-无明显缺陷"), (0.1260, "镜头运动-无明显缺陷"), (0.0322, "色彩-不难看"), (0.1629, "光照-完全准确"), (0.2167, "光照-无明显错误"), (0.0197, "光照-存在"), (0.1360, "光照-极美"), (0.0965, "光照-美"), (0.1549, "光照-不难看"), (0.1294, "首帧形状-完全准确"), (0.0989, "首帧形状-无明显错误"), (0.1884, "首帧形状-不混乱"), (0.1844, "全程形状-完美保持"), (0.2636, "全程形状-不混乱"), (0.1117, "镜头运动-高度动态"), (0.0517, "镜头运动-不微弱"), (0.0256, "物体运动-非常平滑"), (0.4390, "物体运动-完全真实"), (0.2686, "画质-非常稳定"), (0.4293, "细节-非常精细"), (0.0085, "细节-不粗糙"), (0.1276, "细节-不显著粗糙"), (0.0580, "文字-全部正确"), (0.1446, "文字-存在"), (0.3942, "内容-属于物理世界"), ] WEIGHT = np.array([w for w, _ in DIMS]) NAMES = [n for _, n in DIMS] Y, N = 1, -1 # 视频 A:几乎静止的精致画面。prompt 要求「人物走过街道」,它只给了个站着的人。 STATIC = [ N, N, Y, # 指令:没走,但也不算完全没关系 Y, Y, # 构图:精心构图 Y, # 镜头运动无缺陷(没动,自然没缺陷) Y, # 色彩 Y, Y, Y, Y, Y, Y, # 光照:静态画面容易做好 Y, Y, Y, # 首帧形状:完美 Y, Y, # 全程形状:不动=完美保持 N, N, # 镜头动态:没有 Y, Y, # 物体运动:平滑(无运动可挑)、真实(无违反物理) Y, # 画质稳定:不动=极稳 Y, Y, Y, # 细节:静态可以渲染很精细 Y, Y, # 文字 Y, # 属于物理世界 ] # 视频 B:真的走起来了,但运动带来形变、抖动和轻微模糊。 DYNAMIC = [ Y, Y, Y, # 指令:确实照做了 N, Y, # 构图:运动中构图一般 Y, # 镜头运动无明显缺陷 Y, # 色彩 N, Y, Y, N, Y, Y, # 光照:运动中光照没那么完美 Y, Y, Y, # 首帧形状:还行 N, N, # 全程形状:走动中有形变 ← 运动的代价 Y, Y, # 镜头动态:有 Y, Y, # 物体运动:平滑且真实 N, # 画质稳定:运动带来抖动 ← 运动的代价 N, Y, Y, # 细节:动态模糊,不够精细 ← 运动的代价 Y, Y, # 文字 Y, # 属于物理世界 ] def score(answers): """复刻 inference-video.py 的 score()。""" a = np.array(answers) assert len(a) == len(WEIGHT), f"维度不匹配: {len(a)} vs {len(WEIGHT)}" return float(np.mean(a * WEIGHT)) def main(): print("=== 1. 权重结构 ===") print(f" 29 个维度,权重总和 {WEIGHT.sum():.4f},均值 {WEIGHT.mean():.4f}") print(f" 最高 {WEIGHT.max():.4f}({NAMES[int(WEIGHT.argmax())]})") print(f" 最低 {WEIGHT.min():.4f}({NAMES[int(WEIGHT.argmin())]})") print(f" 最高/最低 = {WEIGHT.max() / WEIGHT.min():.0f} 倍") instr = WEIGHT[0] + WEIGHT[1] + WEIGHT[2] print(f"\n 指令遵循三档合计 {instr:.4f},占总权重 {instr / WEIGHT.sum() * 100:.1f}%") aes = WEIGHT[3] + WEIGHT[4] + WEIGHT[6] print(f" 构图+色彩合计 {aes:.4f},占总权重 {aes / WEIGHT.sum() * 100:.1f}%") print("\n=== 2. 静态套利的账 ===") shape_stable = WEIGHT[[16, 17, 22]].sum() motion = WEIGHT[[18, 19]].sum() print(f" 静止能白拿:全程形状保持+不混乱+画质稳定 = {shape_stable:.4f}") print(f" 静止要放弃:镜头高度动态+不微弱 = {motion:.4f}") print(f" 净收益 {shape_stable - motion:+.4f}" f" (yes/no 相差 2 个单位,实际影响 ×2 = {(shape_stable - motion) * 2:+.4f})") print("\n=== 3. 两个视频的实际得分 ===") s_static, s_dynamic = score(STATIC), score(DYNAMIC) print(f" 视频 A(几乎静止的精致画面,没照做 prompt) {s_static:+.4f}") print(f" 视频 B(照做了 prompt,但有形变和抖动) {s_dynamic:+.4f}") winner = "A(静态)" if s_static > s_dynamic else "B(动态)" print(f" 奖励模型更偏好:{winner} 差值 {abs(s_static - s_dynamic):.4f}") print("\n=== 4. 差异来自哪几个维度 ===") diff = (np.array(STATIC) - np.array(DYNAMIC)) * WEIGHT order = np.argsort(-np.abs(diff)) print(" (正值=这一维让静态视频占便宜)") for i in order[:8]: if diff[i] != 0: print(f" {diff[i]:+.4f} {NAMES[i]}") if __name__ == "__main__": main() grpo_minimal.py """GRPO 的最小可运行实现,跑在 VisionReward 的真实权重上。 策略只有一个参数:运动强度 m。这是刻意的——把维度压到 1 维, 才能看清 GRPO 到底把策略推向哪里,而不是被高维噪声淹没。 人为设定的玩具权衡(不是视频物理定律或真实 VLM 的测量): m 越大 → 指令遵循、动态性维度更容易拿 yes m 越大 → 形状保持、画质稳定、细节精细更容易拿 no(运动带来形变/抖动/模糊) prompt 明确要求运动("a person walking across the street"), 所以「照做」在语义上就等于 m 要大。看 GRPO 最后选了什么。 运行 `python grpo_minimal.py`。 """ import numpy as np from reward_hacking import WEIGHT, NAMES # 各维度对运动强度 m 的敏感方向 # +1: m 越大越容易 yes -1: m 越大越容易 no 0: 与 m 无关 SENSITIVITY = np.zeros(29) SENSITIVITY[[0, 1, 2]] = +1.0 # 指令遵循(prompt 要求运动) SENSITIVITY[[18, 19]] = +1.0 # 镜头动态 SENSITIVITY[[20, 21]] = +0.3 # 运动平滑/真实:要有运动才谈得上,但过猛会假 SENSITIVITY[[16, 17]] = -1.0 # 全程形状保持 SENSITIVITY[22] = -1.0 # 画质稳定 SENSITIVITY[[23, 24, 25]] = -0.7 # 细节精细(动态模糊) SENSITIVITY[[7, 8, 10, 11]] = -0.3 # 光照准确/美观:运动中更难保持 def yes_prob(m): """给定运动强度 m∈[0,1],各维度答 yes 的概率。 基线 0.5 表示「说不准」,敏感度把它往两边拉。 """ return np.clip(0.5 + SENSITIVITY * (m - 0.5) * 1.6, 0.02, 0.98) def sample_video(m, rng): """按概率采一个 yes/no 画像,yes=+1 no=-1。""" return np.where(rng.random(29) < yes_prob(m), 1, -1) # 指令遵循的三个维度。把它们置零就得到一个「只看画质」的奖励模型, # 用来对照权重结构对 reward hacking 的影响。 INSTRUCTION_DIMS = [0, 1, 2] def reward(m, rng, weight=None): """复刻 inference-video.py 的 score():np.mean(answers * weight)。""" w = WEIGHT if weight is None else weight return float(np.mean(sample_video(m, rng) * w)) def weight_without_instruction(): w = WEIGHT.copy() w[INSTRUCTION_DIMS] = 0.0 return w def grpo(steps=400, group=32, lr=0.35, sigma=0.15, theta0=0.0, seed=0, weight=None, verbose=False): """单参数 GRPO,返回收敛后的运动强度 m。 每步采 group 个动作,用组内均值/标准差做 baseline 得到 advantage, 再按 advantage 加权更新——不需要 critic,这是 GRPO 的核心简化。 为了看清裸的优化行为,这里省掉了 PPO 式的 clip 和对参考模型的 KL 惩罚。 """ rng = np.random.default_rng(seed) theta = theta0 for step in range(steps): # 策略:m = clip(sigmoid(theta) + noise),高斯探索 mu = 1 / (1 + np.exp(-theta)) eps = rng.normal(0, sigma, group) ms = np.clip(mu + eps, 0.0, 1.0) rs = np.array([reward(m, rng, weight) for m in ms]) # 组内相对优势:减均值除标准差 adv = (rs - rs.mean()) / (rs.std() + 1e-8) # 潜在动作 a = mu + eps 的高斯 score,m=clip(a) 只是环境映射。 # 这不是含边界原子质量的裁剪后 m 分布的逐点 log-density。 # ∂mu/∂theta = mu(1-mu) grad = float(np.mean(adv * eps / sigma ** 2) * mu * (1 - mu)) theta += lr * grad if verbose and (step % 80 == 0 or step == steps - 1): print(f" {step:>4} {mu:.4f} {float(rs.mean()):+.4f}") return 1 / (1 + np.exp(-theta)) def main(): print("=== 1. 优化前后的运动强度 ===") print(" step 运动强度 m 组内平均奖励") m_final = grpo(verbose=True) print(f"\n GRPO 收敛到 m = {m_final:.4f}") verdict = "静态(放弃了 prompt 要求的运动)" if m_final < 0.35 else \ "动态(照做了 prompt)" if m_final > 0.65 else "折中" print(f" → 策略选择:{verdict}") print("\n=== 2. 从不同初值出发是否都收敛到同一处 ===") for name, t0 in (("偏静态 (m≈0.27)", -1.0), ("中性 (m=0.50)", 0.0), ("偏动态 (m≈0.73)", 1.0)): mf = grpo(theta0=t0, seed=1) print(f" 初值 {name:<18} → 收敛 m = {mf:.4f}") print("\n=== 3. 奖励曲面:各运动强度的期望得分 ===") rng = np.random.default_rng(7) w_no_instr = weight_without_instruction() print(" m 完整权重 去掉指令遵循 (每点 4000 次采样)") best, best_no = (-9e9, None), (-9e9, None) for m in np.linspace(0, 1, 11): r = float(np.mean([reward(m, rng) for _ in range(4000)])) r2 = float(np.mean([reward(m, rng, w_no_instr) for _ in range(4000)])) best = max(best, (r, m)) best_no = max(best_no, (r2, m)) print(f" {m:.1f} {r:+.4f} {r2:+.4f}") print(f" → 完整权重的最高点在 m = {best[1]:.1f}({best[0]:+.4f})") print(f" → 去掉指令遵循后最高点在 m = {best_no[1]:.1f}({best_no[0]:+.4f})") print("\n=== 4. 对照:奖励模型不看指令遵循时,GRPO 选什么 ===") print(" step 运动强度 m 组内平均奖励") m_hack = grpo(weight=w_no_instr, seed=0, verbose=True) print(f"\n 收敛到 m = {m_hack:.4f}" f"(完整权重下是 {m_final:.4f})") print(" → 同一个算法、同一个物理约束,只改奖励的权重结构," "\n 策略就从「照做」翻转成「静止不动」。这就是 reward hacking 的来源:" "\n 不是算法坏了,是它精确地最大化了你真正写下来的那个目标。") print("\n=== 5. 两个收敛点在各维度上的差别 ===") p_ok, p_hack = yes_prob(m_final), yes_prob(m_hack) gap = (p_hack - p_ok) * WEIGHT order = np.argsort(np.abs(gap))[::-1] print(" (正值=hack 后的策略在这一维更容易拿 yes)") for i in order[:6]: print(f" {gap[i]:+.4f} {NAMES[i]}" f" (yes 概率 {p_ok[i]:.2f} → {p_hack[i]:.2f})") if __name__ == "__main__": main()
2026年08月23日
7 阅读
0 评论
0 点赞
1
2
粤ICP备2021042327号