Mask-Aware Policy Gradients for Diffusion Language Models
作者: Haran Raajesh, Kulin Shah, Adam Klivans, Philipp Krähenbühl
分类: cs.CL, cs.AI, cs.LG
发布日期: 2026-07-16
备注: Accepted at COLM 2026
💡 一句话要点
提出Mask-Aware策略梯度以提升MDLM的推理能力
🎯 匹配领域: 支柱二:RL算法与架构 (RL & Architecture) 支柱九:具身大模型 (Embodied Foundation Models)
关键词: 掩码扩散语言模型 强化学习 策略梯度 数学推理 编程任务 马尔可夫决策过程 深度学习
📋 核心要点
- 现有方法在处理掩码扩散语言模型时,无法有效估计对数似然,导致推理能力受限。
- 本文提出了一种两阶段动作的马尔可夫决策过程,将生成过程中的令牌放置和位置重新掩码分开处理。
- 通过优化令牌项和掩码项的组合,实验结果在数学推理和编码基准上显著提升,达到新的性能高点。
📝 摘要(中文)
强化学习在提升大型语言模型的推理能力方面已证明有效,但将其扩展到掩码扩散语言模型(MDLM)面临挑战,主要是由于对数似然估计的不可处理性。现有方法仅通过建模令牌预测来近似该对数似然,忽略了生成过程中位置被解掩的顺序。本文观察到MDLM生成在每一步涉及两个决策:在每个掩码位置放置什么令牌以及选择哪些位置重新掩码。我们将此形式化为一个两阶段动作的马尔可夫决策过程(MDP),并展示策略梯度自然分解为令牌项和掩码项。优化这两个项的组合在数学推理和编码基准上取得了最先进的结果,GSM8K得分为87.1%,MBPP得分为53.4%。
🔬 方法详解
问题定义:本文旨在解决在掩码扩散语言模型(MDLM)中进行有效的对数似然估计的问题。现有方法仅关注令牌预测,忽视了生成过程中位置解掩的顺序,导致推理效果不佳。
核心思路:论文提出将MDLM生成过程视为一个两阶段的马尔可夫决策过程(MDP),在每一步中分别决策放置的令牌和重新掩码的位置,从而更全面地捕捉生成过程中的信息。
技术框架:整体架构包括两个主要模块:令牌决策模块和掩码决策模块。令牌决策模块负责选择在掩码位置放置的具体令牌,而掩码决策模块则负责选择哪些位置需要重新掩码。
关键创新:最重要的创新在于将生成过程分解为两个独立的决策过程,使得策略梯度可以自然地分解为令牌项和掩码项。这一方法与现有方法的本质区别在于更全面地考虑了生成过程中的决策。
关键设计:在参数设置上,采用了适应性学习率和特定的损失函数来优化两个决策模块的输出。此外,网络结构设计上,结合了深度学习中的注意力机制,以增强模型对上下文信息的捕捉能力。
🖼️ 关键图片
📊 实验亮点
实验结果显示,优化后的模型在GSM8K和MBPP基准上分别取得了87.1%和53.4%的得分,显著优于现有方法。这表明本文提出的Mask-Aware策略梯度方法在数学推理和编程任务中具有显著的性能提升。
🎯 应用场景
该研究的潜在应用领域包括自然语言处理中的数学推理、编程任务自动化等。通过提升掩码扩散语言模型的推理能力,可以在教育、软件开发和智能助手等多个领域实现更高效的自动化和智能化服务,具有重要的实际价值和未来影响。
📄 摘要(原文)
Reinforcement learning has proven effective for improving reasoning in large language models, but extending it to Masked Diffusion Language Models (MDLMs) remains challenging due to the intractability of the log-likelihood estimation. Existing approaches approximate this log-likelihood by modeling only the token predictions, ignoring the order in which positions are unmasked during generation. We observe that MDLM generation involves two decisions at each step: what tokens to place at each masked position and which positions to remask. We formalize this as a two-stage action MDP, showing that the policy gradient naturally decomposes into a token term and a masking term. Combining optimization of both terms leads to state-of-the-art outcomes on mathematical reasoning and coding benchmarks, with scores of 87.1% on GSM8K and 53.4% on MBPP.