- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
导读
本文围绕 Google Research 仓库中的 Bigger, Better, Faster(BBF)项目,系统讲解如何在 Atari 100K 数据高效基准上,用 JAX 与 Dopamine 框架训练一个高性能深度强化学习智能体。你将掌握仓库的安装步骤、训练命令、BBF.gin配置全参数含义,以及 BBF 核心机制(高回放比率、周期性网络重置与 shrink-and-perturb、SPR 自预测表示、数据增强)在源码中的具体实现与调用链,并能直接复现 SPR、SR-SPR 等其他智能体配置。
一、项目概览:什么是 BBF
bigger_better_faster/README.md 明确指出,本仓库在 JAX 中实现了Bigger, Better, Faster(BBF)智能体,并构建在Dopamine框架之上。BBF 是面向 Atari 100K 数据高效基准设计的一类智能体:它只允许智能体与环境交互约 10 万步(约两个小时的游戏时间),却要求达到接近人类水平的性能,因此"在有限交互预算内最大化样本效率"是它的核心目标。
仓库的一个鲜明特点是配置即算法:SPR(Schwarzer 等,2021)与 SR-SPR(D'Oro 等,2023)并不需要单独的实现,而是直接作为BBFAgent的超参数配置运行。这一点可以从 bbf/train.py 的AGENTS列表得到印证——rainbow、der、dopamine_der、DrQ、OTRainbow、SPR、SR-SPR、BBF共 8 种"智能体"都通过--agent枚举参数选择,最终由同一个create_agent()工厂函数(bbf/train.py)统一创建spr_agent.BBFAgent,区别仅在于加载的 gin 配置文件不同。
二、环境搭建与依赖安装
仓库的安装非常简洁,README 给出的命令是:
pip install -r requirements.txt具体依赖清单见 bigger_better_faster/requirements.txt,其中与运行直接相关的关键依赖包括:
| 依赖 | 版本要求 | 用途 |
|---|---|---|
jax/jaxlib | >= 0.3.14 | 核心数值计算与自动微分后端 |
flax | >= 0.6.3 | 网络定义(spr_networks.py基于 Flaxlinen构建) |
dopamine-rl | == 4.0.5 | 提供JaxDQNAgent、Runner、Atari 环境封装等基础设施 |
gym[atari,accept-rom-license] | <= 0.25.2 | Atari 环境(需接受 ROM 许可) |
ale-py/atari-py | — | Atari Learning Environment 与经典接口 |
gin-config | — | 超参数配置系统(.gin文件解析) |
tensorflow | — | 日志、seed 设置等辅助功能 |
此外,仓库还提供了 bigger_better_faster/bbf/requirements_jax.txt 与 bigger_better_faster/bbf/requirements_long.txt 两个辅助依赖文件,可根据运行环境按需选择。
README 特别提醒:由于 JAX 的安装与操作系统、CUDA 版本强相关,pip install -r requirements.txt可能不足以让 JAX 在 GPU 上正确运行,需要参考 JAX 官方安装说明按平台补充安装步骤。这是该仓库唯一需要读者自行适配外部环境的地方。
三、训练一个 BBF 智能体:入口命令与全部参数
README 给出的本地训练命令如下:
python -m bbf.train \ --agent=BBF \ --gin_files=bbf/configs/BBF.gin \ --base_dir=/tmp/online_rl/bbf \ --run_number=1注意:python -m bbf.train要求以仓库根目录(即包含bbf包的bigger_better_faster上层目录,仓库根路径为bigger_better_faster/)为工作目录运行,因为 bbf/train.py 中的包导入(如from bigger_better_faster.bbf import eval_run_experiment)是基于完整包路径的。
3.1 命令行 Flag 完整说明
bbf/train.py 定义了全部命令行参数:
| Flag | 类型 / 默认值 | 说明 |
|---|---|---|
--agent | 枚举,默认'SPR' | 选择智能体配置,可选rainbow、der、dopamine_der、DrQ、OTRainbow、SPR、SR-SPR、BBF |
--gin_files | 字符串 | gin 配置文件路径列表(如bbf/configs/BBF.gin),必传 |
--base_dir | 字符串,必填 | 训练输出根目录(checkpoint、日志、config.json 写入处) |
--run_number | int,默认 1 | 运行编号,同时作为默认随机种子 |
--agent_seed | int,默认 None | 自定义种子;为 None 时使用run_number |
--no_seeding | bool,默认 True | 为 True 时忽略 run_number,改用int(time.time() * 1e7) % 2**31随机取种 |
--load_replay_dir | 字符串,默认 None | 从固定数据集目录加载初始回放缓冲;None 则不从外部加载 |
--load_replay_number | int,默认 None | 加载固定回放数据时使用的运行编号,默认沿用run_number |
--save_replay | bool,默认 False | 训练结束后将最终回放缓冲保存为固定数据集到${base_dir}/replay_logs |
--data_logging | bool,默认 False | 是否用智能体记录回放缓冲(当前实现直接抛出NotImplementedError) |
--max_episode_eval | bool,默认 True | 是否使用按"固定 episode 数"评估的DataEfficientAtariRunner |
--tag | 字符串,默认 None | 本次运行的标签,会写入 config.json |
3.2 入口执行流程(源码视角)
从 bbf/train.py 的main()可以看出完整启动链路:
- 设置 TensorFlow 行为、GPU 内存增长或隐藏 GPU(当非
run_xm_preprocessing路径时); - 确定随机种子并调用
set_random_seed()(同时设置PYTHONHASHSEED、tf.random、np.random); run_experiment.load_gin_configs()解析 gin 文件与--gin_bindings覆盖项;write_config()将当前 gin 配置、seed、tag、agent 名落盘为base_dir/config.json(bbf/train.py);- 构造
create_agent_fn,默认使用DataEfficientAtariRunner作为 runner; - 启动
jax.profiler.start_server(9999)供性能剖析,随后runner.run_experiment()开始训练。
四、BBF.gin 配置深度解析:一个"配方"看懂全部机制
bbf/configs/BBF.gin 是 BBF 默认配置,也是理解 BBF 算法设计的最佳入口。下面按逻辑分组给出完整参数及其含义。
4.1 基础 DQN 参数(继承自 DopamineJaxDQNAgent)
JaxDQNAgent.gamma = 0.997 JaxDQNAgent.min_replay_history = 2000 JaxDQNAgent.update_period = 1 JaxDQNAgent.target_update_period = 1 JaxDQNAgent.epsilon_train = 0.00 JaxDQNAgent.epsilon_eval = 0.001 JaxDQNAgent.epsilon_decay_period = 2001 JaxDQNAgent.optimizer = 'adam'gamma = 0.997是 BBF 的高折扣因子。配置中另有min_gamma = 0.97,二者配合实现折扣因子的循环退火(详见 4.4)。update_period = 1、target_update_period = 1表示每步环境交互都更新网络,且目标网络通过target_update_tau软更新而非周期性硬拷贝。epsilon_train = 0.0说明训练时几乎完全依赖噪声探索(noisy = False时退化为确定性策略 + 少量评估噪声)。
4.2 Rainbow 风格组件
BBFAgent.noisy = False BBFAgent.dueling = True BBFAgent.double_dqn = True BBFAgent.distributional = True BBFAgent.num_atoms = 51BBF 是一个完整叠加了 Dueling、Double DQN、51 个原子的分布式(C51 风格)价值头、以及(可选)NoisyNet 的 Rainbow 式智能体。BBF 默认关闭 noisy(相比 SPR 配置,SPR 默认noisy = True)。
4.3 核心机制一:高回放比率(Replay Ratio)
BBFAgent.replay_ratio = 64 BBFAgent.batches_to_group = 2 BBFAgent.batch_size = 32replay_ratio = 64意味着每收集 1 个环境转移,就执行 64 次梯度更新,这是 SR-SPR / BBF 打破"回放比率壁垒"的关键设计。在 bbf/agents/spr_agent.py 的set_replay_settings()中可以找到其换算逻辑:
self._num_updates_per_train_step = max(1, self._replay_ratio * self.n_envs // self._batch_size) self.update_period = max(1, self._batch_size // self._replay_ratio * self.n_envs)即每个环境步对应replay_ratio × n_envs // batch_size次更新;batches_to_group将这些更新分批聚合后通过 JIT 的train函数一次执行,从而把梯度计算密集化,充分利用 JAX 的 XLA 编译。
4.4 核心机制二:周期性重置与 shrink-and-perturb
BBFAgent.cycle_steps = 10_000 BBFAgent.reset_every = 20_000 BBFAgent.shrink_perturb_keys = "encoder,transition_model" BBFAgent.shrink_factor = 0.5 BBFAgent.perturb_factor = 0.5 BBFAgent.no_resets_after = 100_000 BBFAgent.max_update_horizon = 10 BBFAgent.update_horizon = 3 BBFAgent.min_gamma = 0.97 BBFAgent.target_update_tau = 0.005 BBFAgent.target_action_selection = True这是 BBF 最具特色的机制——周期性"重启"网络并配合超参数退火:
reset_every = 20_000:每 2 万训练步对网络做一次重置(注释提示:修改回放比率时应同步调整该值);shrink_perturb_keys = "encoder,transition_model":重置时只对编码器与转移模型应用shrink-and-perturb——参数向初始值方向收缩(shrink_factor = 0.5)再叠加扰动(perturb_factor = 0.5),对应源码中jit_reset与interpolate_weights(bbf/agents/spr_agent.py)中old_weight/new_weight的插值实现;no_resets_after = 100_000:训练步数超过 10 万后停止重置(若延长训练需调整);cycle_steps = 10_000:每个重置周期内,update_horizon 从 3 退火到 10、gamma 从 0.997 退火到 0.97。源码中update_horizon_scheduler与gamma_scheduler(bbf/agents/spr_agent.py)使用指数衰减调度器实现这一"由易到难"的循环学习。
4.5 核心机制三:SPR 自预测表示学习
BBFAgent.spr_weight = 5 BBFAgent.jumps = 5 BBFAgent.data_augmentation = True BBFAgent.replay_scheme = 'prioritized' BBFAgent.learning_rate = 0.0001 BBFAgent.encoder_learning_rate = 0.0001jumps = 5:SPR 转移模型在潜空间向前预测 5 步,回放缓冲以subseq_len = jumps + 1的子序列形式采样(见 bbf/agents/spr_agent.py);spr_weight = 5:SPR 辅助损失权重,总损失为loss = dqn_loss + spr_weight * spr_loss,其中spr_loss = ||spr_predictions - spr_targets||²并按轨迹掩码取平均(bbf/agents/spr_agent.py);data_augmentation = True:对观测施加随机裁剪与强度扰动(实现见 bbf/spr_networks.py 的_random_crop、_per_image_random_crop、_intensity_aug);- 学习率上编码器与价值头分离:
encoder_learning_rate与learning_rate各自独立,对应_build_networks_and_optimizer中用optax.masked构造的双优化器(bbf/agents/spr_agent.py)。
4.6 网络结构
BBFAgent.network = @bbf.spr_networks.RainbowDQNNetwork bbf.spr_networks.RainbowDQNNetwork.renormalize = True bbf.spr_networks.RainbowDQNNetwork.hidden_dim = 2048 bbf.spr_networks.RainbowDQNNetwork.encoder_type = "impala" bbf.spr_networks.RainbowDQNNetwork.width_scale = 4 bbf.spr_networks.ImpalaCNN.num_blocks = 2BBF 使用Impala 残差卷积编码器(encoder_type = "impala")+ 2048 隐藏单元 + 4 倍宽度的大网络("Bigger" 的由来)。bbf/spr_networks.py 定义了三种可选编码器:DQN、IMPALA、RESNET。相比之下,SPR 配置使用dqn编码器、hidden_dim = 512、width_scale = 1,可见 BBF 的网络规模显著更大。
4.7 优化器与正则
bbf.agents.spr_agent.create_scaling_optimizer.eps = 0.00015 bbf.agents.spr_agent.create_scaling_optimizer.weight_decay = 0.1优化器参数沿用 DER(van Hasselt 等,2019):Adam 的eps = 1.5e-4,且权重衰减 0.1(这是 BBF 稳定高回放比率训练的重要正则手段;SR-SPR 配置中该项为 0)。
4.8 训练与评估流程参数
DataEfficientAtariRunner.game_name = 'ChopperCommand' atari_lib.create_atari_environment.sticky_actions = False AtariPreprocessing.terminal_on_life_loss = True Runner.num_iterations = 1 Runner.training_steps = 100000 DataEfficientAtariRunner.num_eval_episodes = 100 DataEfficientAtariRunner.num_eval_envs = 100 DataEfficientAtariRunner.num_train_envs = 1 DataEfficientAtariRunner.max_noops = 30 Runner.max_steps_per_episode = 27000training_steps = 100000:即Atari 100K 基准,默认游戏为ChopperCommand(可通过--gin_bindings覆盖,如DataEfficientAtariRunner.game_name='Pong');sticky_actions = False:Atari 100K 基准不使用粘性动作(与人类基准协议一致);terminal_on_life_loss = True:按生命损失截断 episode,是数据高效研究的常见约定;- 评估用100 个并行环境、100 个 episode,训练仅 1 个环境;
max_noops = 30表示每局开始时随机执行最多 30 次空操作。
4.9 回放缓冲
bbf.replay_memory.subsequence_replay_buffer.PrioritizedJaxSubsequenceParallelEnvReplayBuffer.replay_capacity = 200000 bbf.replay_memory.subsequence_replay_buffer.PrioritizedJaxSubsequenceParallelEnvReplayBuffer.n_envs = 1 bbf.replay_memory.subsequence_replay_buffer.JaxSubsequenceParallelEnvReplayBuffer.replay_capacity = 200000 bbf.replay_memory.subsequence_replay_buffer.JaxSubsequenceParallelEnvReplayBuffer.n_envs = 1BBF 使用容量 20 万的子序列回放缓冲(bbf/replay_memory/subsequence_replay_buffer.py),支持多并行环境、以长度jumps + 1的子序列为单位采样(满足 SPR 多步预测需求),并支持prioritized(优先经验回放,基于deterministic_sum_tree)与uniform两种采样方案。
五、扩展配置:SPR、SR-SPR 与其他智能体
仓库的configs目录共提供 8 个 gin 配置:BBF.gin、SPR.gin、SR_SPR.gin、DrQ.gin、OTRainbow.gin、rainbow.gin、der.gin、dopamine_der.gin(见 bigger_better_faster/bbf/configs)。它们与--agent枚举一一对应,运行方式完全相同,只需替换两个参数:
python -m bbf.train \ --agent=SPR \ --gin_files=bbf/configs/SPR.gin \ --base_dir=/tmp/online_rl/spr \ --run_number=1通过对比 bbf/configs/SPR.gin 与 bbf/configs/SR_SPR.gin 可以直观看到"算法即配置"的哲学:
| 维度 | SPR | SR-SPR |
|---|---|---|
回放比率replay_ratio | 64 | 256(更高) |
重置间隔reset_every | 未启用 | 5_000 |
shrink_factor/perturb_factor | — | 0.8 / 0.2 |
batches_to_group | 未设置(默认 1) | 8 |
| 编码器 | dqn,hidden_dim=512,width_scale=1 | cnn,hidden_dim=512,width_scale=1 |
| 噪声 | noisy=True | noisy=False(注释说明 noisy 更慢且损害性能) |
| 权重衰减 | 未设置(默认 0.1 经注释注明参数源自 DER) | 0.0 |
| 更新视界 | 10(固定) | 未显式设置(依赖 reset 周期退火机制) |
| 默认游戏 | Breakout | ChopperCommand |
SPR 对应 2021 年的自预测表示论文,SR-SPR 则是在其基础上加入高回放比率与周期性重置,二者性能与行为差异完全由 gin 参数体现,无需改动任何 Python 代码。
六、评估机制:DataEfficientAtariRunner 与归一化分数
默认--max_episode_eval=True时使用 bbf/eval_run_experiment.py 中的DataEfficientAtariRunner。它与标准 DopamineRunner的关键区别在于:
- 按 episode 数而非步数评估:
_run_eval_phase固定运行num_eval_episodes = 100个 episode,并支持 100 个并行评估环境(num_eval_envs = 100)与one_to_one精确配对模式(bbf/eval_run_experiment.py); - 严格步数上限:训练阶段精确在
training_steps步终止,符合数据高效研究惯例; - 归一化分数:文件内置了 57 个 Atari 游戏的人类/随机得分表(
atari_human_scores/atari_random_scores),并通过normalize_score(ret, game)将原始回报映射到(随机, 人类]区间(bbf/eval_run_experiment.py)。训练中每个 episode 结束都会打印Steps executed / Num episodes / Return / Normalized Return,并同步写入 TensorBoard summary(Train/EpisodeReturn、Eval/NormalizedScore等)。
七、实验结果与消融数据(scores 目录)
仓库 bigger_better_faster/scores 目录存放了 BBF 论文消融实验的原始结果 CSV,按回放比率分组:
- RR2 / RR8 主结果:
RR2_BBF.csv、RR8_BBF.csv; - 消融变体(文件命名即实验标签):
RR2_BBF+sticky_20k.csv~RR2_BBF+sticky_1M.csv:在训练 2 万步至 100 万步区间启用粘性动作的对比;RR2_BBF+γ=0.99.csv:固定折扣因子 0.99(对应去除 gamma 退火);RR2_BBF+n=10.csv:固定更新视界 10(对应去除 update_horizon 退火);RR2_BBF-Annealing.csv:去除周期退火;RR2_BBF-Resets.csv/RR2_BBF-HarderResets.csv:去除重置或使用更激进的重置策略;RR2_BBF-SPR.csv:去除 SPR 自预测损失;RR2_BBF-WD.csv:去除权重衰减。
这些 CSV 与 bbf/configs/BBF.gin 中的开关一一对应,读者可据此验证各机制对最终性能的贡献。注意:仓库仅提供数据文件,未提供绘图脚本,如需可视化需自行读取。
八、实践要点与常见注意事项
- 路径约定:训练命令中的
--gin_files=bbf/configs/BBF.gin是相对路径,需在仓库根目录(含bbf/包的目录,即bigger_better_faster/)下执行; - 环境适配:JAX 的 GPU/CUDA 安装是独立步骤,务必按官方指引完成后再
pip install -r requirements.txt;gym版本被严格锁定为<=0.25.2,且 Atari ROM 需通过accept-rom-license授权; - 随机种子:
--no_seeding默认为 True(随机取种),需要可复现实验时显式传入--no_seeding=False --run_number=N或--agent_seed=N; - 修改回放比率时:
reset_every、batches_to_group需同步调整(gin 注释与set_replay_settings()中的整除断言均提示了这一点,bbf/agents/spr_agent.py); - 延长训练:若
training_steps超过 10 万,需相应增大no_resets_after,否则后期将不再执行重置; - 输出物:
base_dir下会生成config.json(当前 run 的完整 gin 配置快照)、TensorBoard 事件文件与 checkpoints;性能剖析服务默认监听 9999 端口。
参考文献
- Max Schwarzer, Ankesh Anand, Rishab Goel, Devon Hjelm, Aaron Courville and Philip Bachman.Data-efficient reinforcement learning with self-predictive representations. ICLR 2021(对应
SPR.gin配置)。 - Pierluca D'Oro, Max Schwarzer, Evgenii Nikishin, Pierre-Luc Bacon, Marc Bellemare, Aaron Courville.Sample-efficient reinforcement learning by breaking the replay ratio barrier. ICLR 2023(对应
SR_SPR.gin配置)。
以上两篇论文的官方链接可分别从 bigger_better_faster/README.md 的 References 段获取;本仓库实现与其对应的配置、源码路径已在上文各节逐一标注,可直接对照阅读。
- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
相关推荐
Dopamine中的元学习:快速适应新环境的RL算法
Dopamine中的元学习:快速适应新环境的RL算法 引言 在强化学习(Reinforcement Learning, RL)领域,智能体(Agent)通常需要
强化学习机器学习深度学习DreamerV2:基于离散世界模型的强化学习框架技术深度解析
DreamerV2:基于离散世界模型的强化学习框架技术深度解析 在强化学习领域,基于模型的强化学习框架正逐渐成为研究热点。DreamerV2作为这一领域的代表性
AReaL TIR 智能体实战:基于多轮工具调用的数学推理强化学习指南
AReaL TIR 智能体实战:基于多轮工具调用的数学推理强化学习指南 导读 :本文聚焦 AReaL 仓库中 examples/tir 提供的 Tool Int
人工智能大模型强化学习分布式训练AI Agent
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考