
Contrastive RL把对比学习当作目标条件强化学习的 JAX 实现指南【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research本指南以contrastive_rl仓库为对象围绕论文《Contrastive Learning as Goal-Conditioned Reinforcement Learning》arXiv:2206.07568的核心思想展开如何把对比表示学习本身构造为一种强化学习算法使得表示的内积恰好等于目标条件下的价值函数。读完本文你将掌握该仓库的安装方式、lp_contrastive.py的算法与环境切换方法、全部超参数含义以及离线 RL 与图像任务的运行技巧并理解底层 JAX Acme Launchpad 分布式架构与对比学习损失的源码实现。核心思想表示学习本身就是 RL在强化学习中好的状态表示往往能让任务更容易求解。虽然“深度”RL 理应自动学到这样的表示但既往研究常发现端到端学习表示不稳定因此许多算法会额外附加表示学习部件如辅助损失、数据增强。这篇工作的主张恰恰相反不向已有 RL 算法添加表示学习部件而是直接把对比表示学习方法本身改造成一种 RL 算法。具体做法是基于先前工作把对比表示学习应用到带动作标签的轨迹上使得学到的表示的内积恰好对应于一个目标条件下的价值函数。作者由此重新解释了既有 RL 方法其实是在做对比学习并提出一个更简单的算法在大量目标条件 RL 任务上取得更高成功率在基于图像的任务上无需数据增强或辅助目标对比 RL 也优于先前方法。仓库同时提供了论文中的新算法、部分基线算法以及配套环境。若在研究中复用了本仓库建议按 README 给出的 BibTeX 条目引用article{eysenbach2020contrastive, title{Contrastive Learning as Goal-Conditioned Reinforcement Learning}, author{Eysenbach, Benjamin and Zhang, Tianjun and Salakhutdinov, Ruslan and Levine, Sergey}, journal{arXiv preprint arXiv:2206.07568}, year{2022} }仓库结构一览在进入安装与运行之前先了解仓库布局便于后续对照源码lp_contrastive.py实验入口负责环境、算法、计算参数的选择与 Launchpad 程序构建run.sh一键冒烟测试脚本以 debug 模式训练少量步数requirements.txt完整依赖清单锁定版本contrastive/config.pyContrastiveConfig数据类集中定义全部超参数contrastive/agents.pyDistributedContrastive分布式 Agent 定义contrastive/builder.pyLearner/Actor/回放表/数据集迭代器构造器含“以未来状态作为目标”的采样逻辑contrastive/learning.pyContrastiveLearner实现 critic/actor/alpha 三种损失与参数更新contrastive/networks.py策略网络、Q 网络与表示编码器含图像输入分支contrastive/utils.py环境工厂、成功/距离观察者、目标坐标过滤 wrapper 等工具contrastive/distributed_layout.py多 Actor 分布式布局env_utils.py、ant_env.py、fetch_envs.py、point_env.py环境加载与环境实现。环境安装Anaconda Python 3.9README 给出了 5 步安装流程本仓库以requirements.txt锁定依赖版本确保可复现# 1. 获取仓库本镜像仓库已包含全部文件可直接进入 contrastive_rl 目录 # 2. 创建 Python 3.9 环境 conda create -n contrastive_rl python3.9 -y # 3. 激活环境 conda activate contrastive_rl # 4. 安装依赖--no-deps 避免连锁升级 pip install -r requirements.txt --no-deps # 5. 冒烟验证 chmod x run.sh; ./run.sh依赖栈的几个关键点见 requirements.txtJAX 生态jax0.3.13、jaxlib0.3.10、flax0.5.0、dm-haiku0.0.6、optax0.1.2、chex0.1.3RL 框架dm-acme0.4.0、dm-launchpad0.5.0、dm-reverb0.7.0分布式训练与回放缓冲环境依赖gym0.19.0、mujoco2.2.0、mujoco-py2.0.2.7、pybullet3.2.5以及两个 git 依赖d4rl离线 RL 数据集固定 commitd842aa1和metaworld固定 commita0009ed用于 sawyer 系列任务其余tensorflow2.8.2数据集迭代器内部使用tf.data、tensorboard2.8.0、PyOpenGL3.1.6图像渲染。由于部分包通过 git commit 锁定安装时请保证网络可访问对应仓库。--no-deps表示严格按清单逐项安装不自动解析传递依赖有助于复现论文实验环境。运行实验从冒烟测试到完整复现冒烟测试./run.sh实际执行的命令是见 run.shpython lp_contrastive.py --debugTrue--debug标志会强制覆盖一组小型参数见 lp_contrastive.py把训练压缩到数千步内用于验证安装与数据通路是否正常params.update({ min_replay_size: 2_000, local: True, # 本地模式禁用评估器 num_sgd_steps_per_step: 1, # 每步仅 1 次梯度更新 prefetch_size: 1, num_actors: 1, # 单个 Actor batch_size: 32, max_number_of_steps: 10_000, hidden_layer_sizes: (32, 32), # 更小的策略/Q 网络 })完整复现论文实验复现论文结果直接运行环境默认sawyer_windowpython lp_contrastive.pylp_contrastive.py是唯一的实验入口其main函数按三步组织详见 lp_contrastive.py第 1 步选择环境。通过env_name指定支持的家族如下注释中即完整清单Metaworldsawyer_{push,drawer,bin,window}OpenAI Gym Fetchfetch_{reach,push}D4RL AntMazeant_{umaze,medium,large}另有umaze_diverse等变体2D 导航point_{Small,Cross,FourRooms,U,Spiral11x11,Maze11x11}图像观测变体sawyer_image_{...}、fetch_{reach,push}_image、point_image_{...}离线环境offline_ant_{umaze,umaze_diverse,medium_play,medium_diverse,large_play,large_diverse}默认参数块如下env_name sawyer_window params { seed: 0, use_random_actor: True, entropy_coefficient: None if image in env_name else 0.0, env_name: env_name, max_number_of_steps: 1_000_000, use_image_obs: image in env_name, } if ant_ in env_name: params[end_index] 2注意max_number_of_steps的语义分两种在线 RL 时为环境步数离线 RL 时为梯度步数。ant_系列会把end_index设为 2表示只取前两维坐标作为目标。第 2 步选择算法。目前支持五种算法通过修改参数切换算法名参数改动说明contrastive_nce使用默认超参论文主打算法NCE 式对比学习contrastive_cpcuse_cpcTrue改用 CPCInfoNCE 风格目标c_learninguse_tdTrue; twin_qTrue时序差分TD式目标 双 Qncec_learninguse_tdTrue; twin_qTrue; add_mc_to_tdTrue蒙特卡洛项叠加到 TD 目标gcbcuse_gcbcTrue目标条件行为克隆基线文档同时说明许多其他算法可以通过传入其他参数或添加几行代码实现。第 3 步选择计算参数。默认参数已调好主要用于调试如FLAGS.debug分支所示。最终通过get_program(params)构建 Launchpad 程序并启动lp.launch(program, terminalcurrent_terminal)若希望不同组件Actor/Learner/Evaluator分窗口显示可改用terminaltmux。图像实验必须使用多进程基于图像的实验需要 OpenGL 渲染直接运行会遇到 OpenGL 上下文冲突。README 明确要求使用多进程启动方式python lp_contrastive.py --lp_launch_typelocal_mp多线程方式--lp_launch_typelocal_mt亦可使用但图像任务推荐前者。此外lp_contrastive.py 会在检测到use_image_obsTrue且非本地模式时自动覆盖一批参数保证图像任务性能params[num_sgd_steps_per_step] 16 params[prefetch_size] 16 params[num_actors] 10离线 RL 实验的配置要点README 以offline_ant_umaze为例说明离线实验的入口。当env_name以offline_ant开头时程序会做两件事见 lp_contrastive.pynum_actors置 0离线 RL 不需要在线 Actor 收集数据评估另行处理覆盖一组离线专用超参params.update({ samples_per_insert: 1_000_000, # 极大值等效移除限流器 samples_per_insert_tolerance_rate: 100_000_000.0, random_goals: 0.0, # Actor 更新时只用未来状态做目标 bc_coef: 0.05, # 给 Actor 增加行为克隆项 twin_q: True, # 双 Critic取最小值 batch_size: 1024, # 256 → 1024 repr_dim: 16, # 表示维度 64 → 16 hidden_layer_sizes: (1024, 1024), # 策略网络 (256,256) → (1024,1024) })这些调整的源码依据samples_per_insert与容差率共同构造SampleToInsertRatio限流器见 builder.py离线场景用极大值等效关闭bc_coef在 learning.py 中把行为克隆损失与策略梯度损失按bc_coef加权混合twin_q对应 Q 网络最后一维堆叠两个输出并在 actor 更新取最小值见 networks.py 与 learning.py。还有一个环境特例offline_ant_umaze_diverse环境 700 步终止但演示数据有 1000 步因此 lp_contrastive.py 会把max_episode_steps强制设为 1000。超参数速查ContrastiveConfig 全字段所有可调超参集中在 contrastive/config.py 的ContrastiveConfig数据类中分四组损失相关参数默认值含义batch_size256训练批大小actor_learning_rate3e-4策略优化器学习率learning_rate3e-4Q 网络优化器学习率reward_scale1奖励缩放discount0.99折扣因子n_step1n 步回报entropy_coefficientNone熵奖励系数None表示自适应对应 SAC 式 alpha 温度target_entropy0.0目标熵自适应时使用tau0.005目标网络软更新系数hidden_layer_sizes(256, 256)策略与 Q 网络隐层大小回放相关参数默认值含义min_replay_size10000开始采样前的最小回放量max_replay_size1000000回放表最大容量prefetch_size4预取批数num_parallel_calls4数据集并行读取线程数samples_per_insert256采样/插入比限流samples_per_insert_tolerance_rate0.1限流容差率num_sgd_steps_per_step64每步执行的梯度更新次数算法开关参数默认值含义repr_dim64表示维度离线时改 16use_random_actorTrue初始使用均匀随机策略repr_normFalse表示是否 L2 归一化use_cpcFalse使用 CPC 损失localFalse本地模式禁用评估use_tdFalse使用 TD 目标twin_qFalse双 Critic取最小use_gcbcFalse目标条件行为克隆use_image_obsFalse图像观测random_goals0.5Actor 更新中随机目标的比例jitTrueJAX JIT 编译add_mc_to_tdFalse在 TD 目标上叠加 MC 项resample_neg_actionsFalse是否重采样负动作bc_coef0.0行为克隆损失权重环境派生参数obs_dim、max_episode_steps、start_index、end_index由程序根据具体环境自动写入lp_contrastive.py中config.obs_dim obs_dim、config.max_episode_steps ...一般无需手工设置。target_entropy_from_env_spec辅助函数config.py给出了自适应熵的目标值启发式未指定时取-num_actions并要求动作空间为 BoundedArray 且上下界为 ±1——这与lp_contrastive.py中的断言action_spec().minimum -1、maximum 1相呼应。底层原理从源码看对比 RL 如何实现1. 环境构造与目标坐标过滤contrastive/utils.py 的make_environment完成三件事加载 gym 环境、用StepLimitWrapper施加步数上限、再用ObservationFilterWrapper把观测裁剪为[state, goal]拼接形式——其中 goal 取原观测的start_index:end_index坐标。obs_to_goal_2d即对应论文中从状态提取目标坐标的操作。对ant_系列还会套CanonicalSpecWrapper统一 spec 格式。InitialRandomActorutils.py实现“初始随机策略”通过检查策略网络首个线性层偏置是否全零判断是否已被更新未更新时在 [-1, 1] 均匀采样动作更新后切换到策略网络输出。2. 网络结构内积即价值contrastive/networks.py 的make_networks是核心构造器表示编码器_repr_fn对(state, action)拼接后用 MLP 编码为sa_repr对goal用另一个 MLP 编码为g_repr均输出repr_dim维repr_normTrue时做 L2 归一化并可学习温度参数repr_log_scale组合函数_combine_reprjnp.einsum(ik,jk-ij, sa_repr, g_repr)即批内两两内积得到[batch, batch]的 logits 矩阵——这正是“表示内积 目标条件价值函数”的数值载体Q 网络_critic_fn返回上述内积矩阵twin_qTrue时用同一图像表示额外编码一组表示并堆叠成[batch, batch, 2]策略网络_actor_fnMLP NormalTanhDistributiontanh 高斯分布动作输出范围为 [-1, 1]与环境的动作空间断言一致图像输入分支用 Acme 的AtariTorsoCNN 编码器处理 64×64×3 的 state/goal 图像除以 255 归一化。3. 数据集构造未来状态作目标contrastive/builder.py 的make_dataset_iterator用tf.data实现轨迹级采样其关键设计是flatten_fn中的目标采样对每条轨迹以概率正比于discount^(Δt)从未来状态中采样一个作为当前步的 goalis_future_mask * discount加权后tf.random.categorical采样。这正是目标条件 RL 中“以未来状态作为目标”的标准做法也是obs_to_goal_2d在数据管线的落点。数据批构造为[state, goal]与[next_state, goal]的 transition并做了 transpose_shuffle 与随机 shift便于学习器消费。4. Learner 的三种损失contrastive/learning.py 的ContrastiveLearner定义了完整的更新逻辑Critic 损失核心是把[batch, batch]的 logits 与单位矩阵I做对比。MC 风格下use_cpc用 softmax 交叉熵optax.softmax_cross_entropy 0.01 的 logsumexp 平方正则否则用 sigmoid 二分类交叉熵对角线为正样本、其余为负样本。TD 风格下use_td对角元素对应立即下一状态使用w next_v / (1 - next_v)截断到 20对正样本加权并组合loss_pos / loss_neg1 / loss_neg2三项add_mc_to_td时按(1-discount)/((1-discount)1)的比例混合下一状态与原始 goal 作为新目标。更新时按tau软更新目标网络日志中输出binary_accuracy、categorical_accuracy、logits_pos/neg、logsumexp等指标Actor 损失random_goals控制目标重组方式——0.0 用原目标、0.5 拼接原目标与滚动一位的目标、1.0 全用滚动目标损失为alpha * log_prob - diag(q_action)SAC 风格bc_coef 0时叠加行为克隆项Alpha 损失当entropy_coefficientNone时用论文 Eq.18SAC 温度更新见代码注释引用的 arXiv:1812.05905自适应调节熵系数log_alpha初始为 0Adam 学习率 3e-4。每次step()通过process_multiple_batches(update_step, num_sgd_steps_per_step)一次消费多个 batchjitTrue时整体 JIT 编译并统计steps_per_second等日志。5. 分布式布局与评估contrastive/agents.py 的DistributedContrastive继承自 distributed_layout.py 的DistributedLayout是典型的Acme Actor-Learner-Evaluator 分布式程序多个 Actor 并行采集数据写入 Reverb 回放表Learner 从中采样做梯度更新Evaluator 用sample_evalparams.mode()即分布均值评估。Actor 与 Evaluator 均挂载两个观察者utils.pySuccessObserver统计“回合内是否有正奖励”得到success与success_1000指标DistanceObserver测量到目标的 L2 距离输出init_dist / final_dist / delta_dist / min_dist及平滑后窗口统计。评估器默认使用sample_eval确定性模式localTrue时禁用评估。常见问题与提示图像任务报 OpenGL 错误务必使用--lp_launch_typelocal_mp多进程启动离线任务不需要 Actor程序会自动把num_actors置 0评估独立进行offline_ant_umaze_diverse有 1000 步演示数据但环境 700 步终止程序已自动处理--debug只用于安装验证完整复现请直接运行python lp_contrastive.py并按需修改env_name与算法参数动作空间约束程序断言动作范围必须为 [-1, 1]若自定义环境请保持一致问题反馈README 建议直接联系论文作者 Benjamin Eysenbacheysenbachgoogle.com。【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考