- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
导读
本文以 Google Research 仓库中 pse/jumping_task/README.md 为核心,系统讲解 ICLR 2021 论文《Contrastive behavioral similarity embeddings for generalization in reinforcement learning》(PSE)在 Jumping Task(Jumpy World)环境上的完整实验复现流程。你将掌握:如何准备 jumping-task 环境与依赖、如何运行官方冒烟测试脚本、如何用一条命令启动 PSE 训练,以及如何精确复现论文中 wide / narrow / random 三种网格配置与彩色障碍物实验。文章同时深入 train.py、training_helpers.py 等源码,剖析对比损失、表示对齐损失、伪度量定点迭代与 RandConv 等核心机制的底层实现,帮助你理解"行为相似性嵌入"如何在强化学习泛化中发挥作用。
一、实验背景:PSE 与 Jumping Task
PSE(Policy Similarity Embedding)的核心思想是:在强化学习中,与其只依赖奖励或状态特征,不如直接度量"行为相似性"——两个状态如果在策略行为上等价,就应当被映射到相近的表示空间。这一方法通过对比学习(contrastive learning)实现,在训练中显式地把"行为等价的状态"拉近,把"行为不同的状态"推开,从而学习到对任务变化(如障碍物位置、地面高度变化)泛化能力更强的策略表示。
Jumping Task(仓库内称为 Jumpy World)是论文专门用于研究泛化的最小化测试环境:
- 智能体需要在一个不断前进的世界中跨越障碍物;
- 每个环境实例由两个参数决定:障碍物位置(obstacle position)与地面高度(floor height);
- 通过改变这两个参数可以构造出大量相似但不同的任务,天然适合验证"训练任务 → 未见任务"的泛化能力;
- 可选地,障碍物可以被染成不同颜色(红色障碍物行为与白色一致,绿色障碍物行为不同),用于研究视觉外观变化下的泛化。
仓库中该实验的完整实现位于 pse/jumping_task 目录,包含 8 个文件:训练主程序 train.py、数据生成 data_helpers.py、训练辅助函数 training_helpers.py、模型定义 model_helpers.py、评估辅助函数 evaluation_helpers.py、Gym 封装 gym_helpers.py、依赖清单 requirements.txt 以及冒烟测试脚本 run.sh。
二、环境准备
2.1 安装 jumping-task 环境
PSE 的训练代码依赖外部的jumping-taskGym 环境(论文作者维护的独立环境包)。仓库的 run.sh 中给出了标准的安装方式:
git clone https://github.com/Maluuba/jumping-task.git pip install -e jumping-task/gym-jumping-task安装后需要确保环境中存在gym-jumping-task与gym-jumping-colors-task两个注册的 Gym 环境,分别对应白色障碍物与彩色障碍物任务。data_helpers.py 中通过gym_helpers.create_gym_environment('jumping-task')与create_gym_environment('jumping-colors-task', obstacle_color=...)来创建环境。
2.2 安装 Python 依赖
requirements.txt 声明了全部运行依赖:
absl-py gin-config gym matplotlib numpy seaborn tensorflow>=2.0.0注意:
- 代码基于
tensorflow.compat.v2编写(train.py),要求 TensorFlow 2.x; gin-config用于 Gym 环境的配置注入(见 gym_helpers.py 的@gin.configurable);seaborn与matplotlib用于绘制评估网格图并写入 TensorBoard。
三、快速验证:运行冒烟测试
在完整跑实验之前,建议先运行官方提供的冒烟测试脚本,验证代码与环境是否就绪。在仓库根目录(google_research)下执行:
bash pse/jumping_task/run.sh该脚本会依次完成(run.sh):
- 用
virtualenv -p python3 .创建虚拟环境并激活; git clone下载 jumping-task 环境并以 editable 模式安装gym-jumping-task;- 安装
pse/jumping_task/requirements.txt中的依赖; - 以
python3 -m pse.jumping_task.train --train_dir pse --training_epochs 125 --rand_conv启动一次 125 个 epoch、开启 RandConv 的短训练,验证端到端流程可跑通。
这里展示了两种训练入口的差异:
- 在仓库根目录下使用模块路径:
python3 -m pse.jumping_task.train; - 在pse 目录内部使用模块路径:
python -m jumping_task.train(README 的 Launch 部分采用此方式)。
四、启动 PSE 训练:基本命令
从pse目录内部执行:
python -m jumping_task.train --train_dir {TRAIN_DIR} --training_epochs {EPOCHS}两个必选参数(在 train.py 中通过flags.mark_flag_as_required强制校验):
| 参数 | 类型 | 说明 |
|---|---|---|
--train_dir | string | 训练 checkpoint 与 TensorBoard 摘要的保存目录 |
--training_epochs | int | 训练的总 epoch 数(论文复现统一使用 2000) |
训练过程中会做三件重要的事(train.py):
- 每
save_checkpoint_every_n_epochs(默认 40)个 epoch 保存一次 checkpoint; - 每
evaluate_every_n_epochs(默认 20)个 epoch 在全部环境网格上评估一次,把评估网格图写入 TensorBoard; - 在每个 epoch 结束时按
decay_rate(默认 0.999)指数衰减学习率。
五、复现论文主要结果:三种网格配置
论文主结果(RandConv + PSEs)在三种训练任务网格配置上验证,$SEED从 1 取到 100(每个 seed 对应一次独立训练,用于统计均值与方差)。README 给出了三条可直接运行的命令。
5.1 wide 网格(宽松网格)
训练环境覆盖较宽的障碍物位置与地面高度范围:
python -m jumping_task.train --training_epochs 2000 --seed $SEED \ --soft_coupling_temperature 0.01 --alpha 5.0 --temperature 0.5 \ --learning_rate 0.0026 --no_validation --rand_conv5.2 narrow 网格(紧凑网格)
训练环境被压缩到更小的参数区间(障碍物位置 28–38、地面高度 13–17),泛化难度更高:
python -m jumping_task.train --training_epochs 2000 --seed $SEED \ --min_obstacle_position 28 --max_obstacle_position 38 --min_floor_height 13 \ --max_floor_height 17 --positions_train_diff 2 --heights_train_diff 2 \ --soft_coupling_temperature 0.01 --alpha 5.0 --temperature 0.5 \ --learning_rate 0.0026 --no_validation --rand_conv与 wide 网格相比,narrow 网格额外覆盖了 4 个关键参数:
| 参数 | 默认值 | narrow 网格取值 | 含义 |
|---|---|---|---|
--min_obstacle_position | 20 | 28 | 训练可见的最小障碍物位置 |
--max_obstacle_position | 45 | 38 | 训练可见的最大障碍物位置 |
--min_floor_height | 10 | 13 | 训练可见的最小地面高度 |
--max_floor_height | 20 | 17 | 训练可见的最大地面高度 |
--positions_train_diff | 5 | 2 | 障碍物位置的采样步长(决定训练位置密度) |
--heights_train_diff | 5 | 2 | 地面高度的采样步长 |
从 train.py 的参数校验注释可以看到环境参数的合法范围:min_obstacle_position >= 14、max_obstacle_position <= 48、min_floor_height >= 0、max_floor_height <= 41。在 narrow 网格中,训练位置将按步长 2 在[28, 38] × [13, 17]区间上均匀采样,产生约(11/2+1) × (5/2+1) ≈ 42个训练环境组合(实际由 generate_training_positions 以笛卡尔积生成)。
5.3 random 网格(随机网格)
训练环境从完整参数区间内随机采样,检验 PSE 在"随机任务子集"上的学习能力:
python -m jumping_task.train --training_epochs 2000 --seed $SEED \ --soft_coupling_temperature 0.01 --alpha 5.0 --temperature 0.5 \ --learning_rate 0.0026 --no_validation --rand_conv --random_tasks--random_tasks开启后,generate_training_positions 会先根据positions_train_diff/heights_train_diff计算训练环境数量,再用np.random.seed(seed)从完整位置/高度集合中随机抽取相同数量的环境组合,因此不同 seed 对应不同的训练任务集合。
5.4 彩色障碍物实验(PSEs + colors)
为了验证 PSE 对视觉外观变化(颜色)的泛化能力,README 给出了PSEs在彩色障碍物上的复现命令($SEED同样从 1 到 100):
python -m jumping_task.train --training_epochs 2000 --seed $SEED \ --soft_coupling_temperature 0.01 --alpha 5 --l2_reg 0.00007 \ --learning_rate 0.006 --temperature 0.5 --no_validation --use_colors与前面三个命令相比,这一配置:
- 开启
--use_colors:生成红色(RED,行为与白色一致)与绿色(GREEN,行为不同)两种障碍物的模仿数据(见 data_helpers.py 的OBSTACLE_COLORS枚举); - 引入
--l2_reg 0.00007:对网络权重施加 L2 正则(默认 0,即不启用); - 学习率更高(0.006 vs 0.0026)。
当--use_colors开启时,train.py 会使用红色障碍物的数据来生成对比对齐的样本对,而评估阶段则对 WHITE / RED / GREEN 三类环境分别打分并写入 TensorBoard(eval/{split}_{color}_solved等标量,见 train.py)。
5.5 超参数速查表
以上复现命令共同涉及的对比学习超参数定义在 train.py,汇总如下:
| 参数 | 默认值 | 论文复现常用值 | 说明 |
|---|---|---|---|
--alpha | 1.0 | 5.0 | 表示对齐损失(alignment loss)的权重系数 |
--temperature | 1.0 | 0.5 | 对比损失(NT-Xent 风格)中的温度 |
--soft_coupling_temperature | 1.0 | 0.01 | 计算 soft coupling 时的温度 |
--learning_rate | 1e-2 | 0.0026 / 0.006 | Adam 优化器初始学习率 |
--l2_reg | 0.0 | 0.0 / 0.00007 | L2 正则系数 |
--batch_size | 256 | — | 批大小 |
--decay_rate | 0.999 | — | 每 epoch 学习率衰减因子 |
--dropout | 0.0 | — | Dropout 比率 |
--seed | 1 | 1–100 | 随机种子 |
--rand_conv | False | True | 是否在输入上应用 RandConv 随机扰动 |
--no_validation | False | True | 关闭验证集划分(全部环境归入测试) |
--use_coupling_weights | 0 | — | 是否用 coupling 权重加权正负样本 |
--ground_truth_coupling | False | — | 是否使用真实 coupling(仅作对照实验) |
--projection | True | — | 表示学习前是否经过投影层 |
--use_l2_loss | False | — | 是否用 L2 距离损失替代对比损失 |
--use_bisim | False | — | 是否使用 π*-bisimulation 度量替代动作相似性度量 |
其余基线(baseline)与消融(ablation)实验的详细超参数,README 明确指向论文附录 G.3(论文 PDF)。
六、源码级原理剖析
理解了复现命令后,下面从源码角度深入 PSE 在 Jumping Task 上的实现机制。
6.1 整体训练流程
main → train_agent → train_step 构成核心训练链路:
- 构建
JumpyWorldNetwork策略网络(2 个动作:跳/不跳); - 生成全网格的模仿学习数据(专家轨迹);
- 按网格配置挑选训练环境位置集合;
- 每个训练 step:计算交叉熵损失 + 可选的 L2 正则 + 可选的表示对齐损失,加权求和后反传更新。
train_step中损失组合的核心逻辑(train.py):
total_loss += cross_entropy_loss # 模仿学习主损失 total_loss += l2_reg * l2_regularization_loss # 仅当 l2_reg > 0 total_loss += alpha * alignment_loss # 仅当 alpha > 0即 PSE 的训练目标 = 行为克隆损失 + α × 表示对齐损失(+ 可选的 L2 正则)。
6.2 表示对齐损失与伪度量
对齐损失的核心实现在 training_helpers.py 的representation_alignment_loss:
- 从训练环境中随机抽取一对环境(sample_train_pair_index),取出两者的观测、专家动作与奖励;
- 用网络提取两组观测的表示,计算两两间的余弦相似度矩阵(默认)或 L2 距离矩阵;
- 计算两组专家轨迹之间的代价矩阵:默认用动作相似性(
calculate_action_cost_matrix),若开启--use_bisim则改用奖励差(calculate_reward_cost_matrix); - 用定点迭代求解满足贝尔曼风格的伪度量
d_metric(metric_fixed_point):
d_metric_new[i, j] = action_cost_matrix[i, j] + \ gamma * d_metric[min(i + 1, n - 1), min(j + 1, m - 1)]该迭代以gamma=0.999(见 train.py)逐步扩散,"行为越相似、未来行为也越相似"的轨迹对会获得越小的度量值,即期望的伪度量目标;
- 最后把相似度矩阵与该伪度量送入对比损失,使表示空间的相似度逼近行为相似性度量。
6.3 带 soft coupling 的对比损失
contrastive_loss 实现了带 soft coupling 的对比损失:
- 相似度矩阵先除以温度
temperature; - 以伪度量矩阵的逐行/逐列
argmin确定正样本对; - 若启用 coupling 权重(
use_coupling_weights),则用coupling = exp(-metric_values / coupling_temperature)构造软耦合矩阵,把负样本按耦合强度加权(这也是--soft_coupling_temperature 0.01的作用——温度越小,coupling 越接近硬匹配); - 损失为对称形式
loss1 + loss2(分别按行、按列聚合 log-sum-exp)。
6.4 RandConv:输入随机化
RandConv(model_helpers.py)来自网络随机化方法:对输入施加一个权重每次前向时随机重采样的卷积层。核心逻辑:
- 以概率
alpha=0.1跳过扰动(原样返回输入); - 否则用
glorot_normal重新初始化卷积核并作用于输入,模拟随机的视觉扰动。
在 train.py 中,开启--rand_conv后训练前会先调用一次rand_output预热层结构;评估时 evaluation_helpers.py 采用蒙特卡洛平均(eval_mc_samples = 5,见 train.py)——对同一观测做 5 次前向、平均 softmax 输出后再取贪心动作,以抵消随机卷积的方差。
6.5 网络结构
JumpyWorldNetwork 是一个小型卷积策略网络:
conv0: Conv2D(32, 8×8, stride 4)→conv1: Conv2D(64, 4×4, stride 2)→conv2: Conv2D(64, 3×3, stride 1);flatten → dense0(256) → dense1(64) → dense2(2);representation()方法输出 64 维表示(projection开启时经过dense1),供对齐损失计算相似度;call()则在表示后接 dropout 与动作输出层,用于策略预测。
6.6 数据生成与评估
- 专家轨迹生成:generate_optimal_trajectory 针对每个(位置, 高度)组合,在障碍物前
JUMP_DISTANCE=13步处执行"跳"动作(绿色障碍物除外),并断言累计奖励以保证轨迹最优;观测通过 stack_obs 堆叠相邻帧形成 2 帧输入。 - 训练数据组织:create_balanced_dataset 将跳/不跳两类样本分开,用
sample_from_datasets以 1:1 比例上采样少数类,缓解模仿学习中的类别不平衡。 - 评估:create_evaluation_grid 在全部环境组合上计算贪心动作与专家动作的差异,判断任务是否被解决;num_solved_tasks 按 train / validation / test 三个划分分别统计"已解决任务数",写入 TensorBoard 的
eval/{split}_solved标量,作为泛化性能的量化指标。
七、训练产物与结果解读
训练完成后,--train_dir下会生成两类产物(train.py):
model/:checkpoint(默认最多保留 1 份,可用--max_checkpoints_to_keep调整);tb_log/:TensorBoard 摘要,包括各损失曲线(loss/total_loss、loss/cross_entropy_loss、loss/alignment_loss、loss/l2_regularization_loss)、学习率曲线,以及评估网格热力图(Grid/Evaluation/{color})与解决任务数标量(eval/{split}_solved)。
此外,开启--show_alignment_loss_image后,还会记录对齐损失中的 coupling 代价矩阵与相似度矩阵图像(align/coupling_cost、align/similarity_matrix),便于可视化表示是否正确地与行为相似性对齐(train.py)。
使用 TensorBoard 观察:
tensorboard --logdir {TRAIN_DIR}/tb_log八、复现注意事项
- 模块入口差异:仓库根目录下用
python -m pse.jumping_task.train,pse 目录内用python -m jumping_task.train,两者等价,注意当前工作目录与sys.path的关系。 - 随机种子:
--seed同时作用于 TensorFlow、NumPy 与PYTHONHASHSEED(set_random_seed),复现论文统计量需要从 1 到 100 完整跑一遍并聚合结果。 - 环境参数合法性:障碍物位置与地面高度的取值受环境限制(
>= 14 / <= 48、>= 0 / <= 41,见 train.py),narrow 网格中的参数均在合法范围内。 - 验证集开关:论文复现命令统一使用
--no_validation(验证位置为空列表,train.py),全部环境作为测试集衡量泛化;若需观察验证曲线,去掉该 flag 后 train.py 会根据网格是否紧凑自动生成验证位置。 - 依赖版本:代码使用
tensorflow.compat.v2与tf.kerasAPI,需要 TensorFlow 2.x;gym-jumping-task需先于训练脚本安装,否则import gym_jumping_task(data_helpers.py)会失败。
- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
相关推荐
Contrastive RL:把对比学习当作目标条件强化学习的 JAX 实现指南
Contrastive RL:把对比学习当作目标条件强化学习的 JAX 实现指南 本指南以 contrastive_rl 仓库为对象,围绕论文《Contrast
人工智能深度学习NLP计算机视觉强化学习DeepSeek-R1对比学习:相似性推理应用
DeepSeek R1对比学习:相似性推理应用 引言:新一代推理模型的突破 在人工智能快速发展的今天,大型语言模型(LLM,Large Language Mod
基础模型大模型人工智能DeepSeekVLM-R1与SFT方法对比:为什么强化学习在跨域泛化上更胜一筹
VLM R1与SFT方法对比:为什么强化学习在跨域泛化上更胜一筹 在当今多模态AI快速发展的时代,视觉语言模型(VLM)的训练方法选择直接影响着模型的性能表现。
人工智能大模型多模态计算机视觉强化学习微调模型评测Ascend
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考