跳转至

Why Deterministic PRM Guidance Underperforms in Discrete Diffusion Reasoning

会议: NeurIPS2026
arXiv: 2609.35472
代码: https://github.com/dLLM-PRM-Gap/
领域: LLM 推理
关键词: 离散扩散语言模型、过程奖励模型、测试时计算、候选多样性、结果验证器

一句话总结

这篇诊断研究把扩散生成和奖励评分统一计入推理预算,发现 Dream-7B 的确定性 top-1 PRM 引导不如独立采样加任务匹配 ORM,并通过候选池、终局评分和读出消融拆出失败发生的位置。

研究背景与动机

离散扩散语言模型不是逐个固定左到右的前缀,而是反复去噪一个部分被遮蔽的完整序列。一个中间快照可能已经露出后面的答案片段,却还缺少前面的运算条件。这样的生成过程看似适合过程奖励模型(PRM):既然解答尚未定型,就可以检查中间状态,把后续计算分给更有希望的分支。然而,自回归推理中有效的前缀验证器,并不天然适合这种散落在不同位置的可见证据。

问题还不只是奖励模型的分类能力。引导必须额外调用评分器,并且会提前删除候选;一个早期高分但最终错误的状态,可能让其他正确解答的祖先永久消失。另一方面,擅长区分不同题目总体难易的评分器,不一定能在同一道题的多个答案中选对。因此,单独报告 PRM 的 ROC-AUC,或者只看引导后的正确率,都不足以说明额外计算到底花得是否值得。

本文没有提出一个新的生成网络,而是对已有的分段分支与 top-1 剪枝配方做成本受控的诊断。核心 idea:把“搜索后还剩哪些正确候选”与“评分器是否选中正确候选”分开测量,再用终局专用验证器、保多样性的采样和读出消融定位瓶颈。

方法详解

整体框架

输入是数学题或编程任务,生成器以部分遮蔽的解答序列为状态,经过 128 次去噪得到完整答案。诊断首先使用训练题产生带结果标签的快照并拟合验证器,再按统一的前向调用预算比较独立采样与引导搜索,最后分别评估候选池和最终选择。

这套流程中的 cross-mask PRM 是“用最终结果监督的中间状态价值模型”,不是用人工逐步推理正确性标签训练的传统过程监督模型。它在不同遮蔽率下预测该状态所属轨迹的最终正确性;ORM 和 final-state PRM 则只在完全解码的状态上训练。前者负责评估中间状态,后两者专门评估完整答案,训练数据分布与推理职责不能混为一谈。

下图是诊断流程,不是论文提出的新网络;虚线表示训练监督或消融关系,实线表示评测数据流。PRM Guided 输出经过终局剪枝的一个答案,PRM Hybrid 则保留终局的全部候选供候选池诊断。

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
    A["训练题轨迹<br/>最终正确性标签"] -.-> B["结果监督快照"]
    B -.->|拟合 PRM 与终局验证器| C["统一成本协议"]
    Q["测试题"] --> C
    C -->|PRM 引导或独立采样| D["候选池与选择分离"]
    D -->|Hybrid 池或独立池| P["Oracle 上限与多样性"]
    D -->|Guided 终局 PRM 或 ORM 选择| O["答案正确率"]
    D -.-> E["读出对照"]
    E -.-> O

关键设计

1. 结果监督快照:先明确 PRM 实际学到的是什么

作者在 GSM8K 训练题上运行生成器,保存去噪快照,并把一条完整轨迹的最终正确或错误标签分配给它的中间状态。这种监督不需要人工为每个步骤打标签,但其含义是“这次续写最后是否成功”,而不是“当前已出现的每一步都正确”。同一部分状态仍可能被不同随机续写补成正确或错误答案,尤其在高遮蔽率下,继承标签只是状态价值的带噪观测。

PRM 使用冻结的扩散语言模型骨干、可训练 LoRA 和两层 MLP 奖励头,将解答区域的池化隐藏表示与 256 维去噪步索引嵌入结合起来。主比较保持数据与奖励头一致,只改变双向注意力或因果注意力;读出策略另外做消融。终局专用模型只看完整解答,final-state PRM 还移除步索引嵌入,因此“PRM”与“ORM”的名字本身并不能解释性能差距,关键是实际训练协议。

高遮蔽率时快照缺少判别最终结果的证据,双向 PRM 的 ROC-AUC 从接近完成的遮蔽桶约 0.77 降到最高遮蔽桶约 0.54。附录 F.1 的信息论上界说明:在固定遮蔽桶和类别先验下,如果状态与最终正确性的互信息趋近于零,任何评分器的最优 AUC 上限也趋近随机水平。但这是一条条件性结论,并没有证明真实去噪分布中的互信息必然随遮蔽率单调变化,也不意味着当前 PRM 已达到理论上限。

2. 统一成本协议:把中间评分和终局评分都算进预算

PRM Guided 从全遮蔽状态开始,每隔 \(b\) 个去噪步复制当前状态为 \(K\) 个分支,让各分支独立采样,再把 PRM 分数最高的一个留下。下一段的全部分支都从这个共同祖先继续,而不是从头独立生成。最后一段仍然需要给所有候选评分并做 top-1 选择;当 \(b\) 不整除步数时,最后一段缩短。

记总去噪步数为 \(T\),独立完整样本数为 \(N\),计算量以一个候选通过一个 dLLM 规模模型的一次前向为单位:

\[ C_{\mathrm{ORM}}(N)=NT+N,\qquad C_{\mathrm{PRM}}(K,b)=KT+K\left\lceil\frac{T}{b}\right\rceil. \]

ORM Rerank 把预算用于 \(N\) 条独立完整轨迹,再对每条做一次终局评分。主设置 \(T=128,b=64\) 下,PRM Guided 是 \(130K\) 次前向,ORM Rerank 是 \(129N\) 次;令 \(K=N\) 只是约 0.8% 以内的近匹配,不是严格相等。微小预算差异对 PRM 有利,也不能把这种前向匹配直接写成墙钟时间相等,因为模型调用、并行和分段开销不同。训练验证器的计算另计,不包含在这些推理成本公式中。

3. 候选池与选择分离:正确答案不见了,还是没有选中它

PRM Hybrid 与 Guided 的前面各段相同,但最后一段不删除候选,而是保留全部 \(K\) 个终局答案。这样才能计算 Oracle@K:只要池中至少有一个正确答案,就把该题记为可解。Oracle 使用真实正确性,属于理想选择器的上限,不是可部署的算法,也不应把 Hybrid 的池上限写成 Guided 的实际正确率。

作者同时统计每题的不同答案数,并做离线反事实剪枝:在存储的早、中、晚状态上按分数保留 top-M,看剪枝是否删除了所有最终能成功的谱系。保留两个状态时,每个状态下一段生成 \(K/2\) 个子候选,因此总分支宽度仍是 \(K\)。这控制了“放宽剪枝只是偷偷多花计算”的解释。

另一个对照是同预算的有效样本量(ESS)调温 SMC:保留多个粒子,以调温后的 PRM 分数赋权,必要时系统重采样。SMC 明显恢复候选池,但仍用 cross-mask PRM 做最终加权答案投票或选最高分粒子。若 Oracle 恢复而实际正确率不恢复,问题就不止在多样性,还在最终选择。对独立的同一候选池分别使用 cross-mask PRM、ORM 和 final-state PRM,则进一步隔离终局评分能力。

Pooled ROC-AUC 是跨所有题目混合后的正确/错误候选区分能力,并不等于题内排序能力。一个评分器只会给简单题高分、困难题低分,也可能有很高的 pooled AUC,却给同题候选相同分数。本文的 PRM 不是完全没有题内信号,但中等的成对排序优势不足以保证最高分候选一定正确;附录的独立比较假设只是示意模型,不是实际 PRM 正确率的通用上界。

4. 读出对照:避免把池化缺陷全归咎于因果注意力

因果模型的早期 token 隐藏状态没有读取完整后文;对所有可见 token 做均值池化,会混入大量只见过局部信息的表示。作者保持其他训练设置不变,把读出改为最后一个非 MASK、非 EOS 位置的隐藏状态。完全解码时,这个位置可以汇总前面的完整文本,因此能直接检查主因果 PRM 的失败是否来自不合适的读出。

最后 token 池化把终局 ROC-AUC 从 0.6147 提到 0.7274,回收了相对双向均值池化 0.7759 的约 70% 差距。它也改善重排序,但 \(N=32\) 时仍只有 50.64%,远低于 ORM 的 82.71%;修好读出并不等于解决跨遮蔽训练造成的终局专用能力差距。更长训练与终局训练对照仍留下双向优势,但不能据此断言所有因果序列级评分器在理论上都无法处理快照;附录 F.4 只覆盖受限的前缀读出。

一个完整示例

以 \(K=8,b=64,T=128\) 为例,第一次生成八个半完成快照需要 512 次生成前向和八次 PRM 评分。若 top-1 留下的快照含有错误运算,另外七条谱系之后即使可能给出正确答案,也不能再回来。第二段再从这个保留状态产生八个完整答案,花费同样的 520 次前向。

Guided 随后只输出终局 PRM 最高分的答案;Hybrid 把这八个完整答案都交给诊断。如果八个都错,任何终局选择器都无法补救;如果有正确答案却没有选中,才是终局评分问题。相同规模的 ORM Rerank 则花 1,032 次前向生成八条独立完整解答并终局评分,避免中途把全部候选绑在一个祖先上。这里的调用数是协议计算,不是一个新测得的案例成绩。

原文还有一个终局误排序实例:新增设备每天耗电 2 kWh,电价为每 kWh 1.50 美元,每周新增费用应是 21 美元。三个独立候选中,错误答案 30 美元的因果 PRM 分数为 +1.001,错误答案 15 美元为 −0.827,正确答案 21 美元反而为 −0.908;ORM 从完整的 32 候选池选中了正确解答。这说明“池里有正确答案”和“验证器选择正确”不同,但不是实测某条 Guided 轨迹全过程的重建。原文用 \(t=13,25,3\) 标记这三个候选,此处将其视为候选标识,不误读成去噪步数。

损失函数 / 训练策略

主 PRM 用最终正确性二元标签和二元交叉熵(BCE)训练,LoRA 作用于 q_proj 与 v_proj,rank 为 16、alpha 为 32。主比较训练 2,000 步,batch size 为 32,学习率按余弦从 \(2\times10^{-5}\) 降到零。骨干冻结;训练、调参和早停只用训练题,验证题与拟合题分离,不使用测试题训练。

GSM8K 主采样设置为 temperature=0.5、alg_temp=0.5、top_p=1.0;训练快照和测试使用相同采样器与快照计划。更长训练的 15,000/31,000 步,以及终局协议下的 8,407 步,属于不同对照,不能替换成主 PRM 的训练预算。

为检验继承标签噪声,作者还对 10,000 个训练状态各做八次新续写,用成功比例作为分数标签,在相同的 500 步训练控制下比较。这将 pooled AUC 从 0.816 提到 0.827,但推理正确率只从 60.88% 到 61.49%,差异的区间包含零。它说明更细的结果价值监督在这个控制中没有消除主要差距,不等于人工过程标签、更多续写或其他监督一定无效。

实验关键数据

主实验

主证据来自 Dream-v0-Instruct-7B。GSM8K 使用全部 1,319 个测试题和严格答案提取;32 候选池共 42,208 条完整轨迹。下表只摘取主结果与任务匹配控制,正确率单位为 %,前向次数是每题推理成本。

任务与预算 PRM Guided ORM Rerank ORM 优势(百分点) 成本与口径
GSM8K,K=N=8 65.18 ± 0.75 75.13 9.95 1,040 vs 1,032;Guided 三次运行均值
GSM8K,K=N=32 70.02 ± 0.70 82.71 12.69 4,160 vs 4,128;Guided 两次运行均值
MATH,K=N=8 20.80 30.65 9.85 任务专用验证器,500 个测试题
MBPP,任务匹配控制 50.88 ± 3.57 63.04 ± 0.84 12.16 257 个留出任务;主文未在该行单列 N/K

来源为原文表 1、表 2、附录 B.2 与 D.1。MATH 验证器拟合/验证分别用 450/50 题,与留出的 500 个测试题分开;同一存储候选池重评分的 ORM 数字是 31.90%,不是表中 30.65%,不能拼接成一个结果。MBPP 的 PRM 终局重排序为 65.47 ± 0.74%,与 ORM 的差异区间包含零,却明显高于 Guided 的 50.88%,使它成为“终局 PRM 已经不弱,但引导仍有损失”的控制。

GSM8K 的单样本基线为 43.14%,多数投票在 N=8/32 时为 60.05%/67.63%。因此 Guided 仍有实际信号并超过多数投票,只是同等规模计算用来独立采样加终局 ORM 更划算。ORM@8 还超过所有已测 Guided 预算,包括四倍规模的 K=32,以及 Guided 扫描最高的 K=24、73.67%。

消融实验

GSM8K,K=N=8 的搜索方式 Oracle@8(%) 每题不同答案数 cross-mask PRM 最终选择(%) ORM 最终选择(%)
独立采样 81.05 4.31 42.84 75.13
top-1 引导的 Hybrid 池 67.30 ± 1.24 1.75 65.18(Guided 输出) 未运行
同预算 SMC 77.89 ± 0.57 3.95 ± 0.02 65.48 ± 0.12(加权投票) 未运行

来源为原文表 3、附录 B.1 与 D.2;Hybrid 候选池和 Guided 最终输出是配套诊断,但不是把两者视为同一个输出接口。SMC 最高分粒子的正确率另为 66.34 ± 0.82%。top-1 的 Oracle 损失是 13.75 个百分点,SMC 回收 10.59 个百分点,却没有把实际选择提高到 ORM 水平,支持候选池损伤和终局选择错误是两个不同瓶颈。表中“未运行”不能被理解成 ORM 在恢复池上无效。

终局评分器 完全解码状态 pooled ROC-AUC N=8 重排序正确率(%) N=32 重排序正确率(%)
因果 cross-mask PRM,均值池化 0.6147 43.90 40.56
因果 cross-mask PRM,最后 token 池化 0.7274 49.66 50.64
双向 cross-mask PRM,均值池化 0.7759 42.84 65.35
双向 ORM,终局专用训练 0.9623 75.13 82.71
Final-state PRM,终局专用训练 未单列 75.40 ± 0.05 82.79 ± 0.11

来源为附录 C.3、C.4 与 B.2,候选为相同的独立采样终局池。Final-state PRM 与 ORM 的配对差异区间均包含零;表格不把 ORM 的 AUC 当成 final-state PRM 已报告的 AUC。双向 cross-mask PRM 在 N=8 时甚至低于因果最后 token 版本,说明 pooled AUC 的相对排序也不能直接推出某个预算下的题内 top-1 效果。

关键发现

  • 放宽 top-1 为 top-2,在同一分支宽度下把 Guided 从 65.18% 提到 69.70%,仍比 ORM@8 低 5.43 个百分点;离线末状态剪枝删光正确谱系的风险,top-1/top-4 分别是 19.92%/5.93%。
  • 分支间隔 b=16/32/48/64 的单次评测为 59.4%/62.8%/65.1%/66.5%。66.5% 是单次运行,而主表 65.18% 是多次均值;附录 C.2 的单次比较差距 8.79/11.98 也不能替代主表的 9.95/12.69。
  • 近完成快照桶的 AUC 为 0.7702,完全解码终局为 0.7759;0.816/0.827 属于新续写监督控制。它们来自不同切片,不能混成同一条遮蔽率曲线。
  • 墙钟验证中 Guided@8 约 156 秒/题,ORM@8 约 168 秒/题,并不严格时间匹配;ORM@6 已能在约 126 秒达到 72.40%,仍超过 Guided@8 的 65.18%。

亮点与洞察

  • 把 Oracle 上限与实际正确率并列,可以判断奖励搜索是在“删掉答案”还是“看错答案”。它比只报告最终性能更能指导改进:前者需要保谱系,后者需要改终局验证器。
  • Final-state PRM 与 ORM 相当,是对模型命名误导的直接纠正。不能因为一个模型叫 PRM,就假定它必须持续指导搜索或必然不适合终局验证,实际训练分布更重要。
  • SMC 对照说明恢复多样性是必要方向,但不是这套实验中的充分条件。未来可测试保多样性的搜索搭配终局专用评分器;这是后续假设,不是本文已经验证的组合成绩。

局限与展望

  • 结论针对 Dream-7B 上的确定性分段 top-1 引导及所测控制,不覆盖所有 PRM、所有扩散语言模型或所有随机搜索算法;任务也限于数学与代码,开放文本生成未测。
  • LLaDA 的跨骨干证据是双向/因果 PRM 引导的配置均值 31.64%/22.25%,单样本为 20.77%。没有 LLaDA 专用 ORM 对照,不能把它写成第二个骨干也验证了 ORM 全面胜出;标准差是跨配置而非同配置随机种子方差。
  • GSM8K 训练的 ORM 不迁移到 MATH500:N=32 时只有 6.10%,但该 OOD 表的独立采样与 Guided 温度设置不同,只能支撑迁移失败,不能代替任务匹配 MATH 控制。
  • 去掉右半生成 token 后 AUC 仅变约 0.001,反转输入的因果评分近随机,但这没有排除所有右侧证据作用。反转文本还改变了骨干熟悉的顺序与 RoPE 相对位置,不能当成因果架构在原则上无能的证明。
  • 自适应晚期引导、显式多样性核、追加 CLS 或学习查询读出、人工步骤监督均值得验证。主结果不包含验证器训练摊销;原文整体研究约用 2,400 H20 GPU-hours,不能拿推理调用近匹配声称总训练加推理成本相等。

相关工作与启发

  • vs 自回归 PRM 与 Math-Shepherd:它们主要评估按顺序增长的推理前缀;本文面对非前缀的遮蔽快照,并使用结果继承标签而非人工逐步标注。启发是先检查监督含义和状态分布,再迁移奖励搜索配方。
  • vs 自一致性与 best-of-N 验证:多数投票主要依靠答案重复频率,ORM 使用训练出的终局正确性信号。本文显示强任务匹配验证器的收益不能简单归为“多采样几次”。
  • vs dLLM 粒子搜索、重遮蔽与 reward-free guidance:这些工作改变采样或构建替代奖励;本文提供成本和池/选择分解,适合用作共同评测协议,但没有实测所有这些方法孰优孰劣。

评分

  • 新颖性: 4/5 — 新意在失败分解与受控诊断,不是新生成器。
  • 实验充分度: 4/5 — 同池、SMC、读出、监督与任务对照较完整,跨骨干 ORM 及更广搜索仍缺。
  • 写作质量: 4/5 — 因果链条清楚,但单次/多次、数据切片及不同协议需仔细区分。
  • 价值: 4/5 — 为奖励引导研究提供可复用的成本与候选池验收标准。