【源码解读版】ContraBAR:用 CPC 对比学习替代变分推断做贝叶斯元 RL

文章目录

导读

项目地址:https://github.com/ec2604/ContraBAR

论文解读:https://blog.csdn.net/m0_59012280/article/details/164032246?spm=1001.2014.3001.5501

ContraBAR (Contrastive Bayes-Adaptive Reinforcement learning,对比贝叶斯自适应强化学习) 是一种元强化学习框架,它通过对比预测编码发现任务表示,并利用这些表示来条件化策略,从而在各种不同任务间实现快速适应。

与重建过往经验(如 VariBAD 等变分方法)不同,ContraBAR 通过将当前信念状态与来自不同任务的负样本进行对比,学习区分智能体正面临哪个任务------从而为贝叶斯自适应深度强化学习提供了一种更具样本效率和可扩展性的方法。

在标准强化学习中,智能体针对单一固定任务进行训练。而在元强化学习中,智能体面临的是一个任务分布------每个任务都有自身的动力学或奖励结构------并且必须学会在少数几轮交互内进行适应。

核心挑战在于任务推断 :智能体必须仅利用目前收集到的轨迹,迅速弄清自己正处于哪个任务中。ContraBAR 通过训练 RNN 编码器,经由对比损失(而非重建损失)来生成任务信念表示,从而解决了这一问题。随后,该信念会被输入到一个基于当前状态进行决策的标准策略网络(PPO/A2C)中。

bash 复制代码
ContraBAR/
├── main.py                    ├── metalearner_cpc.py         # 带有 CPC 编码器的核心元训练循环
├── learner.py                 # 基线学习器(无元学习 / 先知)
├── cpc.py                     # CPC 模块:编码器初始化、损失计算、负采样
├── models/
│   ├── encoder.py             # RNNCPCEncoder, RNNEncoder, ImageEncoder
│   ├── policy.py              # Actor-Critic 策略网络
│   └── cpc_modules.py         # MLP 分类器, actionGRU, statePredictor
├── algorithms/
│   ├── ppo.py                 # PPO 实现
│   ├── a2c.py                 # A2C 实现
│   ├── sac.py                 # SAC 实现
│   └── online_storage.py      # 用于策略更新的采样缓冲区
├── environments/
│   ├── wrappers.py            # contrabarWrapper (MDP → BAMDP)
│   ├── parallel_envs.py       # 向量化环境创建
│   ├── navigation/            # PointRobot, GridWorld
│   ├── mujoco/                # Ant, Cheetah, Walker, Humanoid 变体
│   ├── dm_control/            # Reacher
│   └── panda_gym/             # PandaReacher
├── config/                    # 各环境专属超参数配置
├── utils/
│   ├── storage_cpc.py         # 用于 CPC 的零填充采样存储
│   ├── helpers.py             # 实用函数(动作选择、编码)
│   ├── evaluation.py          # 评估与可视化
│   └── tb_logger.py           # TensorBoard / WandB 日志记录
└── figures/                   # 论文图表与架构图

快速开始

ContraBAR 基于 PyTorch 构建,使用 OpenAI Gym 环境,并使用 TensorboardX(可选 Weights & Biases)进行日志记录。

相比于同类型项目代码,多了一个 Kornia 库,用于图像增强(像素环境)

本仓库提供了三个 conda 环境文件及匹配的 pip 依赖文件,对应于你的目标环境所需的 MuJoCo 版本:

文件 MuJoCo 版本 适用场景
contrabar_200.yml / requirements_200.txt mujoco 2.3.1 默认选择 --- PointRobot、GridWorld、dm_control Reacher、Panda 以及所有基于图像的环境
contrabar_150.yml / requirements_150.txt mujoco 1.50 Cheetah、Ant(方向/目标)运动
contrabar_131.yml / requirements_131.txt mujoco 1.31 Walker、Hopper 运动
  1. 克隆并创建环境(使用推荐的默认版本 MuJoCo 2.x):
bash 复制代码
git clone https://github.com/ec2604/ContraBAR.git
cd ContraBAR
conda env create -f contrabar_200.yml
conda activate contrabar

contrabar_200.yml:项目提供的环境描述文件,里面写死了这个项目依赖的 python 版本、cuda 版本、pip 包、conda 包(pytorch、torchvision、各类科学计算库)

  1. 安装 pip 依赖(用于 yml 中未包含的特定环境包):
bash 复制代码
Copy code
pip install -r requirements_200.txt

和其他项目代码一样 :ContraBAR 使用单一入口点 main.py , 配合 --env-type 标志从 config/ 目录加载相应配置。随后,训练流水线会实例化 MetaLearner(它协调 CPC 编码器和 RL 策略)并调用其 train() 方法。

启动 PointRobot 实验(最简单的连续导航任务,无 MuJoCo 依赖):

bash 复制代码
python main.py --env-type pointrobot_contrabar

结果默认保存至 ./logs(可通过 --results_log_dir /path/to/dir 配置),且 Tensorboard 事件文件会自动写入。若配置了 Weights & Biases,运行记录也会同步至 varibad_cpc 项目。

参数 默认值 (PointRobot) 作用
--num_processes 16 并行环境数;若 GPU 内存有限则减小
--num_frames 8e7 总训练帧数;为快速实验可减小(例如 1e6)
--lr_policy 4e-5 策略学习率
--lr_representation_learner 1e-4 CPC 编码器学习率
--negative_factor 15 CPC 损失中每个正样本对应的负样本数
--latent_dim 50 信念/隐空间的维度
--seed 73 接受一个或多个种子用于多种子运行
--policy ppo RL 算法:ppo 或 a2c

评估与可视化

训练完成后,使用提供的脚本评估并可视化训练好的策略:

评估每回合回报 --- 计算跨回合的测试时适应性能:

bash 复制代码
python eval_test_perf.py

渲染并可视化 Agent 行为 --- 加载保存的模型,并使用编码器的信念状态渲染 rollout:

bash 复制代码
python viz_script.py

评估特定检查点 --- eval_script.py 从日志目录加载保存的 encoder.ptpolicy.pt 并运行评估回合:

bash 复制代码
model_location = './logs_SparsePointEnv-v0/your_run_folder/'
encoder = torch.load(model_location + 'models/encoder.pt')
policy = torch.load(model_location + 'models/policy.pt')

预计算的基准结果以 .npy 文件形式存储在 end_performance_per_episode/ 中,比较了 ContraBAR 与 VariBAD 及循环基线在六个 MuJoCo 环境上的表现。
README 记录了 MuJoCo 2.0 中的一个已知 Bug:对于 AntGoal 环境,80% 的环境状态会变为零。请为 ant_goal_contrabar 使用 MuJoCo 1.50(requirements_150.txt / contrabar_150.yml)。基于图像的变体(ant_goal_image_contrabar)在 MuJoCo 2.0+ 下可正常运行。

项目框架概述

双模式调度

main.py 中的入口点根据 --env-type 选择配置,并分派到两种训练模式之一:

模式 CPC 编码器 隐变量到策略 用例
元学习 MetaLearner ✅ contrabarCPC 从经验中进行任务推断
基础学习 Learner ❌ 无 先验 / 平均性能基线

当 args.disable_metalearner 为 False(默认值)时,完整的 MetaLearner 路径被激活。精简版的 Learner 完全移除了编码器、隐变量和 CPC 缓冲区,作为非元学习控制组。

MetaLearner --- 训练协调器

MetaLearner 类是中央协调器。其 train() 方法每次迭代实现一个三阶段循环:

参考:metalearner_cpc.py

  • 重新编码:通过 CPC 编码器对整个运行中的轨迹进行重新编码,以获取当前隐藏状态(因为自上次前向传播以来编码器权重可能已更改)。
  • rollout:策略在环境中执行 policy_num_steps 步,每步在线更新编码器的隐藏状态,并将转移数据同时插入 CPC 缓冲区和策略缓冲区。
  • 更新:如果 iter < pretrain_len,则仅预训练 CPC 编码器;否则,通过 PPO/A2C 联合更新策略,并通过对比损失更新 CPC 编码器。

contrabarCPC --- 对比任务推断

contrabarCPC 类封装了所有表示学习逻辑。它包含:

参考:cpc.py

  • RNNCPCEncoder --- 基于 GRU 的历史编码器,将 (action, state, reward) 序列映射到隐藏状态。
  • actionGRU(可选)--- 辅助 GRU,基于未来动作对 CPC 分类器进行条件化,用于多步前瞻预测。
  • MLP 分类器 --- 每个前瞻因子对应一个分类器,用于预测编码后的转移是正样本(同一任务)还是负样本(不同任务)。
  • RolloutStorage --- 专用于已完成轨迹的独立回放缓冲区,在 CPC 更新时独立于同策略策略缓冲区进行采样。

compute_cpc_loss() 方法对一小批次轨迹进行编码,采样负样本(跨任务),构建正/负预测对,并返回交叉熵损失。update_cpc() 方法每次调用执行 num_representation_learner_updates 次梯度更新。

RNNCPCEncoder --- 历史编码

编码器通过三阶段流水线处理交互历史:

参考:models/encoder.py

每种输入模态(动作、状态、奖励)通过一个 FeatureExtractor(线性层 + 激活函数)进行投影。对于基于图像的观测,ImageEncoder(带有谱范数卷积的 CNN → FC → LayerNorm → tanh)将替代状态的 FeatureExtractor。

拼接后的嵌入依次通过可选的 GRU 前全连接层、LayerNorm、GRU 单元本身,以及可选的 GRU 后全连接层。与 VAE 方法中使用的 RNNEncoder 不同,CPC 编码器输出确定性隐藏状态------不应用重参数化技巧或采样。

Policy --- Actor-Critic 网络

参考:models/policy.py

Policy 网络是一个双头 Actor-Critic 网络。它接受状态、隐变量(CPC 隐藏状态)和任务(先验任务,若可用)的任意组合作为输入,每种输入可选择性地通过专用编码器进行归一化和嵌入。拼接后的表示输入独立的 Actor 和 Critic MLP 堆栈。Actor 输出 DiagGaussian(连续)或 Categorical(离散)动作分布的参数。

CPC 模块 --- 分类器和动作 GRU

参考: models/cpc_modules.py

MLP 分类器应用 LayerNorm → Linear → ELU → Linear,为每个(正、负)样本对生成一个标量 logit。actionGRU 是一个输入端带 LayerNorm 的辅助 GRU,用于基于未来动作序列对 CPC 预测进行条件化,以实现时间前瞻。statePredictor 是一个三层回归网络,由可选的表示质量评估器使用。

环境集成

contrabarWrapper --- BAMDP 构建

参考: environments/wrappers.py

contrabarWrapper 通过将视野扩展至 episodes_per_task × H + (episodes_per_task - 1),将单回合 MDP 转换为贝叶自适应 MDP(BAMDP)。在单个 BAMDP 回合内,底层 MDP 多次重置而任务保持固定,使智能体有多次尝试来识别并利用该任务。当 add_done_info=True 时,该包装器在观测中追加一个 done_mdp 位,使状态在回合边界间满足马尔可夫性。

并行向量化环境

参考:environments/parallel_envs.py

双缓冲区架构

ContraBAR 维护两个独立的存储系统,服务于不同的学习目标:

参考:utils/storage_cpc.py, algorithms/online_storage.py

缓冲区 用途 数据流 更新触发
CPC RolloutStorage RolloutStorage 对比表示学习 已完成轨迹 (CPU),运行中轨迹 (GPU) cpc_encoder.update_cpc()
策略 CPCOnlineStorage CPCOnlineStorage 同策略强化学习 (PPO/A2C) 逐步附带隐藏状态的转移数据 policy.policy_update()

CPC 缓冲区在 CPU 上存储完整的零填充轨迹,通过随机添加阈值(representation_learner_buffer_add_thresh)控制回放。策略缓冲区在 GPU 上存储固定长度的同策略 rollout,并附带用于自举的隐藏状态。

训练阶段

元训练循环包含三个由迭代次数和帧数控制的独立时间阶段:

阶段 条件 CPC 更新 策略更新 行为
预收集 frames < precollect_len 随机 rollout 以填充缓冲区
预训练 iter < pretrain_len 仅编码器的对比学习
联合 iter ≥ pretrain_len 包含双重更新的完整元强化学习
相关推荐
白色机械键盘16 分钟前
《垂直领域大模型落地“实践指南“:选型、微调与评估》—— 基于 LLaMA-Factory + Easy-Dataset 的全链路实战
人工智能
现代野蛮人19 分钟前
【深度学习实验】—— 基于 LSTM 的年度医疗费用回归预测
深度学习·回归·lstm
RisunJan20 分钟前
鸿蒙(HarmonyOS NEXT)开发小白入门学习计划表
学习·华为·harmonyos
PM老周22 分钟前
2026年 AI 研发管理工具怎么选?需求、计划、风险与知识四项能力
人工智能·项目管理·研发管理
m0_5474866624 分钟前
《深度学习技术基础》全套PPT课件2026
人工智能·深度学习·powerpoint
征尘bjajmd29 分钟前
黑马AI大模型机器学习课程笔记(个人记录、仅供参考)
人工智能·笔记·机器学习
计科杨某人32 分钟前
从复杂系统到大语言模型:理解 AI 中的“涌现”
人工智能·语言模型·自然语言处理·ai涌现
GFDAGDS32 分钟前
从课程设置看近屿智能AI培训:直播教学、AI互动学习与项目实战如何结合
人工智能·大模型应用·近屿智能
蒸蒸yyyyzwd36 分钟前
cpp选手备战秋招学习笔记day17
笔记·学习