基于掩码的预测表示用于强化学习
Mask-based Predictive Representations for Reinforcement Learning
📝 TLDR
视觉深度强化学习需要从高维图像输入中提取有效状态以提升样本效率。本文提出一种基于掩码预测的自监督辅助任务,利用智能体收集的序列信息预测被掩码内容,结合 Transformer 在隐空间中重建被掩码的输入序列。该方法将学到的压缩表征输入强化学习模型,在多个连续和离散控制基准上显著提升了样本效率,并超越了当前最优方法。该工作验证了非重建式掩码预测作为强化学习辅助任务的有效性。
🧭 速览
视觉深度强化学习面临高维图像输入和样本稀缺的挑战,亟需有效表征以提升样本效率。
提出基于掩码预测的自监督辅助任务,利用序列上下文预测被掩码信息,结合 Transformer 在隐空间重建序列。
在多个连续和离散控制基准上超越当前最优的样本高效强化学习方法,样本效率显著提升。
验证了非重建式掩码预测作为强化学习辅助任务的有效性,为样本高效视觉强化学习提供新思路。
📊 论文图表(共 7 张)
展开查看 7 张图
TL;DR
视觉深度强化学习长期受困于高维图像输入带来的样本效率问题。本文提出一种基于掩码预测的自监督辅助任务(MPR),让智能体在潜空间中预测被掩码的连续帧信息,而非像 MAE 那样重建像素空间。该方法在 DeepMind Control Suite 和 Atari-100k 两大基准上显著超越了当前最先进的样本高效强化学习方法,验证了非重建式掩码预测作为辅助任务的巨大潜力。
研究背景与动机
视觉深度强化学习的核心挑战在于如何从高维像素观测中提取出对决策真正有用的状态表示。与低维向量状态不同,图像包含大量与任务无关的冗余信息——背景纹理、光照变化、相机噪点——这些信息若被不加区分地编码,不仅会大幅增加计算开销,更会干扰智能体捕捉环境的关键动态。传统的像素级重建方法(如自编码器)虽然能压缩维度,但其优化目标与强化学习的策略提升目标往往不一致,学到的表示未必对决策最优。
已有的预训练方法试图通过大规模无标签数据学习通用视觉表征,再迁移到强化学习中微调。然而这类方法存在两个根本性缺陷:一是预训练数据分布与目标任务环境往往存在差异,学到的表示未必适配特定任务;二是预训练与在线学习分离,无法随着智能体在环境中交互而动态调整表示。近年来自然语言处理领域的 BERT 和 GPT 通过掩码语言建模展示了自监督学习的强大能力,计算机视觉领域的 MAE 和 I-JEPA 则证明了掩码预测在视觉表征学习中的有效性。这些成功启发作者思考:能否将掩码预测思想引入强化学习,作为一种与策略学习联合优化的辅助任务?
关键洞察在于,智能体在环境中收集的轨迹数据本身就蕴含丰富的时序结构——相邻帧之间高度相关,智能体的动作会影响未来观测的变化。这种序列上下文信息为预测被掩码的内容提供了天然线索,而掩码预测任务则迫使编码器学习更具预测性、更能捕捉环境动态的表征。
方法
MPR 的整体架构借鉴了 I-JEPA 的联合嵌入预测架构思想,采用在线-目标双编码器结构。给定连续 帧原始观测序列 ,系统首先生成两个视图:上下文视图 和目标视图 。上下文视图经过随机图像增强后,再在时序上应用掩码操作;目标视图则不做掩码,直接作为预测的参照。
掩码策略是方法的关键设计之一。作者采用分块掩码策略,将序列划分为若干块,每块内的帧共享相同的掩码模式。具体而言,若块大小 ,则只在单帧内掩码(空间掩码),适合 Cheetah run 这类以空间特征为主的任务;若 ,则整条序列被同等掩码(时间掩码),适合 Finger spin 这类需要捕捉动作序列模式的任务;介于二者之间则为时空掩码。掩码区域以随机中心放置矩形块实现,掩码率固定在 40% 左右。这一设计使智能体必须同时利用空间和时间维度的上下文信息来推断被遮盖的内容。
编码器层面,在线编码器 负责提取上下文视图的特征,经 Transformer 预测解码器 和投影头 处理后得到预测潜表示 。目标编码器 则以指数移动平均的方式缓慢更新:
这种不对称更新机制确保目标表示保持稳定,为在线表示提供一致的优化目标。预测解码器采用两层多头自注意力机制,结合可学习位置编码,在潜空间中重建被掩码的部分。
损失函数的设计体现了方法的核心洞察——在潜空间而非像素空间进行对齐:
余弦相似度损失只要求预测与目标方向一致,无需精确匹配数值,从而避免了像素级重构带来的信息冗余问题。总体优化目标为强化学习损失与掩码预测损失的加权组合:
推理阶段仅使用训练好的目标编码器输出表征输入策略网络,上下文编码器与解码器退出工作,计算开销与普通强化学习相当。
实验与结果
实验在两类基准上系统验证了方法的有效性。连续控制任务采用 DeepMind Control Suite,选取 6 个代表性环境(Finger spin、Cartpole swingup、Reacher easy、Cheetah run、Walker walk、Ball in cup catch),分别评估 100k 和 500k 环境步下的性能。离散控制任务采用 Atari-100k 基准,包含 26 款游戏,每款游戏限定 100k 步交互(对应约 400k 实际环境步)。
在 DMControl-100k 基准上,MPR 在 6 个任务中的 5 个取得领先,平均得分 830.0,较此前最优的 MLR 方法提升 7.5%,中位数得分 870 提升 4.1%。在 DMControl-500k 条件下,优势略有收窄但仍保持领先,均值达到 910.1。值得注意的是,在困难任务 Reacher hard 和 Walker run 上,MPR 相比 MLR 分别取得约 130 分的显著提升,显示出方法在复杂场景中的适应性。
Atari-100k 结果进一步印证了方法的通用性。在 26 款游戏中,MPR 在 11 款上取得所有方法中的最佳成绩,另有 5 款游戏(Boxing、Freeway、Jamesbond、Krull、Road Runner)超过了人类平均水平。尽管在 Pong、Gopher 等游戏中表现不及 MLR,但整体聚合指标显示 MPR 具备竞争力。
消融实验揭示了多个重要发现。掩码率 40% 是连续和离散任务的共同最优选择,过高会丢失过多上下文信息,过低则无法有效促使编码器学习预测性表征。预测解码器深度 2 层足以实现良好性能,增加深度带来的增益有限但计算成本显著上升。序列长度方面, 在两类任务上均为最优,过长反而引入冗余时序信息干扰表示学习。
讨论与可借鉴点
MPR 的成功揭示了一个重要规律:在强化学习中,表征学习的目标应当是捕捉环境的预测性结构,而非精确重建观测细节。像素级重构要求编码器保留所有视觉信息,其中大量与决策无关的冗余内容会干扰策略优化;而潜空间掩码预测只需编码高层语义和动态规律,这与强化学习的决策需求天然对齐。余弦相似度损失优于 MSE 损失的经验发现进一步支持了这一观点——它表明智能体需要的不是精确重建,而是对环境状态的一致性理解。
从工程角度看,该方法的通用性令人印象深刻。它可以无缝接入 SAC、Rainbow 等主流强化学习算法,连续与离散动作空间均适用,且推理阶段不引入额外计算开销。这为实际部署提供了良好基础——训练时可获得表征学习的好处,部署时仅需执行标准策略推理。
当前方法的局限同样值得深思。掩码策略仍基于简单的中心矩形块,缺乏语义层面的考量——不同任务可能需要掩码不同语义区域才能获得最佳表征。针对每个任务手动调优块大小 、掩码率等超参数的成本较高,缺少自适应机制。Atari 实验直接基于 MLR 代码库进行,可能未能充分挖掘方法潜力。此外,论文未披露实验算力信息,给复现工作带来不确定性。
未来方向包括:探索更具语义的掩码策略(如基于显著区域或物体检测的掩码)、研究自适应掩码机制、以及将方法与大规模视觉基础模型结合。理论层面,深入理解潜空间预测为何优于像素重构、为何余弦相似度优于 MSE,将有助于设计更有效的表征学习目标。
摘要
基于视觉的深度强化学习需要处理高维的图像信息输入。如何从高维图像输入和有限样本中抽象出有效的状态,对于样本高效的强化学习至关重要。为应对这一挑战,受自然语言处理和计算机视觉等领域的启发,我们提出了一种基于掩码预测的自监督任务,将其作为强化学习的辅助任务。这种非重构方法利用智能体从环境中收集的序列信息以及序列中的上下文信息来预测被掩码的内容,从而增强智能体对任务的理解并学习有效的表示。结合 Transformer,我们发现该模型在潜空间中重构了被掩码的输入序列。通过将这种方法学到的压缩表示输入到强化学习模型中,我们观察到强化学习样本效率的提升。此外,该模型在多个连续和离散控制基准测试上优于目前最先进的样本高效强化学习方法。
Abstract
Vision-based deep reinforcement learning involves dealing with high-dimensional inputs of image information. It is crucial to abstract effective states from high-dimensional image inputs and limited samples for sample-efficient reinforcement learning. To address this challenge, inspired by fields such as natural language processing and computer vision, we propose a self-supervised task based on mask prediction as an auxiliary task for reinforcement learning. This non-reconstruction method uses the sequence information collected by the agent from the environment and the context information in the sequence to predict the masked information, thereby strengthening the agent's understanding of the task and learning effective representations. Combined with transformers, we find that the model reconstructs the masked input sequence in the latent space. By feeding the compressed representations learned by this method into reinforcement learning models, we observe an improvement in the sample efficiency of reinforcement learning. Moreover, the model outperforms state-of-the-art sample-efficient reinforcement learning methods on multiple continuous and discrete control benchmarks.
论文详细总结(自动生成)
论文总结:基于掩码的预测表示用于强化学习(MPR)
一、核心问题与研究动机
视觉驱动的深度强化学习(DRL)需直接从高维像素输入中学习策略。然而,真实环境中的样本采集代价高昂且存在安全风险,因而样本高效强化学习成为关键挑战。具体问题包含两方面:
- 高维冗余:图像观测空间维度高、信息冗余大,若直接在像素空间重构表征(如 MAE 类方法)将带来巨大计算开销。
- 预训练—在线学习脱节:已有预训练方法(如 MVP 等)需提前收集数据,无法与智能体在线策略更新同步,难以适应特定任务需求。
受 NLP(BERT、GPT)和 CV(MAE、I-JEPA)中掩码建模成功的启发,作者将"掩码预测"思想引入 RL 作为辅助自监督任务,通过让智能体在潜空间中预测被掩码的连续帧信息,获得更具预测性、一致性的状态表示。
二、方法论
2.1 整体框架
MPR 基于联合嵌入预测架构(JEPA)思想,采用"在线—目标"双编码器结构与 Transformer 解码器,整体流程如图 1 所示:
1. 对原始连续观测序列 进行图像增强与掩码,生成上下文序列 与目标序列 。
2. 上下文序列经在线编码器 编码后,再经 Transformer 预测解码器 与投影头 ,得到预测潜表示 。
3. 目标序列经动量编码器 得到目标潜表示 。
4. 通过余弦相似度损失在潜空间中对齐二者(而非重构像素空间)。
2.2 关键技术细节
- 掩码方式:在每张图像中心随机放置 像素的掩码块(DMControl),Atari 中掩码大小 ,掩码率约 40%。
- 分块掩码策略:将长度为 的序列划分为若干块,块内采用相同掩码;块大小 为"空间掩码", 为"时间掩码",介于二者之间为"时空掩码"( 控制视角重复次数)。
- 目标编码器更新:采用 EMA 方式:
- 预测解码器:两层多头自注意力(MHSA),深度 ,单头注意力,结合可学习位置编码。
- 损失函数:余弦相似度损失
- 联合优化:
- 推理阶段:仅使用训练好的动量目标编码器输出表征给策略网络,上下文编码器与解码器停止工作。
2.3 算法流程(伪代码要点)
算法 1 描述了完整训练流程:交互采样→序列掩码→增强编码→预测解码→计算 与 →联合反向传播更新在线参数→EMA 更新目标参数。
三、实验设计
3.1 基准与环境
- 连续控制:DeepMind Control Suite(DMControl),选取 6 个任务(Finger spin、Cartpole swingup、Reacher easy、Cheetah run、Walker walk、Ball in cup catch),分别在 100k 与 500k 环境步下评估(DMControl-100k / DMControl-500k)。
- 离散控制:Atari-100k 基准,26 款游戏,每局 100k 步(实际交互 400k 步)。
3.2 基线与对比方法
- 连续控制:PlaNet、Dreamer、SLAC、CURL、DrQ、SAC+AE、MLR 等。
- 离散控制:SimPLe、CURL、DrQ、SPR、MLR、CoIT 等。
- 强化学习基座:连续控制用 SAC;离散控制用 Rainbow + MLR(直接套用 MPR 掩码策略,)。
3.3 关键结果
- DMControl-100k:MPR 在 5/6 任务上优于 MLR,均值 830.0(提升 7.5%),中位数 870(提升 4.1%)。
- DMControl-500k:5/6 任务领先,均值 910.1(提升 1.5%)。
- Atari-100k:26 款游戏中 11 款取得最佳,5 款超过人类水平(Boxing、Freeway、Jamesbond、Krull、Road Runner)。
- 困难任务(Reacher hard、Walker run):在 Reacher hard-500k 上较 MLR 平均提升 130 分。
四、资源与算力
论文未明确披露所使用 GPU 型号、数量、单次训练时长及总算力消耗等细节。仅在附录列出了:
- DMControl 实验超参数(学习率 、批大小 512 等);
- Atari 实验超参数(批大小 128、Adam 优化器、100k 训练步等)。
无法从正文判断实验计算开销和复现成本,是一个信息披露不足之处。
五、实验数量与充分性
实验设置较为充分,主要包括:
- 基准对比:DMControl 100k/500k(6 任务 × 10 种子)+ Atari-100k(26 游戏 × 10 种子 × 100 集),规模较大;
- 消融实验:掩码比例(10%–90%)、预测解码器深度(1/2/4/8)、序列长度(4/8/16/24)、分块大小与掩码策略(1/2/4/8)、替代设计(动作 token、、MSE 损失、特征掩码等);
- 困难任务泛化测试:Reacher hard、Walker run;
- 统计方式:连续控制报告均值±标准差,离散控制 10 种子平均,方法客观性较高。
总体实验设计系统、对比方法覆盖 SOTA,但 Atari 部分与 MLR 对比时直接套用了对方代码库而未做大幅改动,可能存在因代码库差异带来的潜在偏差。
六、主要结论与发现
1. 掩码预测辅助任务可显著提升视觉 RL 样本效率,在 DMControl 与 Atari 两大基准上均优于 SOTA。
2. 潜空间预测优于像素空间重构:避免了高维像素重构带来的计算开销,也避免了冗余像素信息的干扰。
3. 序列上下文信息对预测至关重要:相邻帧高度相关,利用连续帧信息预测被掩码内容能强化时空维度上的表征一致性。
4. 最优超参组合:掩码率 40%、解码器深度 2、序列长度 8、(DMControl)/ (Atari)。
5. 掩码策略因任务而异:Finger spin、Walker walk、Ball in cup catch 偏时间掩码;Cheetah run 偏空间掩码;Cartpole swingup 与 Reacher easy 偏时空掩码。
6. 余弦相似度损失优于 MSE:说明无需精确重构潜表示。
7. 与 MLR 的差异:(i) 掩码不依赖 patch 尺寸;(ii) 连续控制中不引入动作序列信息;(iii) 仅一个投影预测器,参数更少。
七、优点与亮点
- 方法设计创新:将 NLP/CV 领域成熟的掩码建模思想成功迁移到 RL,潜空间预测而非像素重构,兼顾效率与效果。
- 通用性强:可与 SAC、Rainbow 等多种 RL 基座结合,连续/离散动作空间均适用。
- 系统消融:对掩码率、解码器深度、序列长度、掩码策略、损失函数、 值、动作注入等多维度均做了消融分析。
- 实验充分:覆盖主流基准且统计严谨(多种子、多集数、报告标准差)。
- 实用价值:降低了视觉 RL 对样本量的依赖,为现实世界部署(如机器人)提供了样本高效方案。
- 未来方向明确:作者指出可结合 SAM、BLIP 等模型生成更具语义的实例级或图文掩码,为后续工作提供了清晰路径。
八、不足与局限
- 算力信息披露缺失:未报告 GPU 型号、数量、训练时长,无法评估实际计算成本与可复现性。
- 超参敏感且任务依赖:分块大小、掩码率、 等需针对不同任务调优,缺乏统一自适应方案,部署成本较高。
- 掩码方式简单:采用随机中心矩形掩码,未探索更具语义信息的掩码(如显著区域、物体掩码),提升空间可能未充分挖掘。
- Atari 实验改动有限:直接在 MLR 代码基础上替换掩码策略,未重新调优,可能未充分发挥方法潜力。
- 应用范围受限:在 Atari 部分仅报告聚合分数,缺少定性分析(如学习到的表征可视化、注意力图)。
- 理论分析欠缺:对"为何潜空间预测优于像素重构"以及"为何余弦相似度优于 MSE"等现象仅给出经验性解释,缺乏理论支撑。
- 未评估鲁棒性与安全性:未在噪声观测、对抗扰动等环境下测试方法的稳健性,限制了其在真实场景中的可信度。
- Atari 中 26 款游戏中仅 11 款取得最佳,且 Pong、Gopher 等游戏反而落后于 MLR,说明方法并非在所有任务上均稳定优越。
(完)
✨ 编译论文
点「✨ 编译」开始,LLM 会按 Polaris 风格翻译并把图片/表格嵌到对应位置。结果存到浏览器 localStorage,下次访问自动加载。






