十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

Bigger Better Faster(BBF):基于 JAX 与 Dopamine 的 Atari 100K 数据高效强化学习智能体实战指南

Bigger Better Faster(BBF):基于 JAX 与 Dopamine 的 Atari 100K 数据高效强化学习智能体实战指南 人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载导读本文围绕 Google Research 仓库中的 Bigger, Better, FasterBBF项目系统讲解如何在 Atari 100K 数据高效基准上用 JAX 与 Dopamine 框架训练一个高性能深度强化学习智能体。你将掌握仓库的安装步骤、训练命令、BBF.gin配置全参数含义以及 BBF 核心机制高回放比率、周期性网络重置与 shrink-and-perturb、SPR 自预测表示、数据增强在源码中的具体实现与调用链并能直接复现 SPR、SR-SPR 等其他智能体配置。一、项目概览什么是 BBFbigger_better_faster/README.md 明确指出本仓库在 JAX 中实现了Bigger, Better, FasterBBF智能体并构建在Dopamine框架之上。BBF 是面向 Atari 100K 数据高效基准设计的一类智能体它只允许智能体与环境交互约 10 万步约两个小时的游戏时间却要求达到接近人类水平的性能因此在有限交互预算内最大化样本效率是它的核心目标。仓库的一个鲜明特点是配置即算法SPRSchwarzer 等2021与 SR-SPRDOro 等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.2Atari 环境需接受 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 \ --agentBBF \ --gin_filesbbf/configs/BBF.gin \ --base_dir/tmp/online_rl/bbf \ --run_number1注意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_numberint默认 1运行编号同时作为默认随机种子--agent_seedint默认 None自定义种子为 None 时使用run_number--no_seedingbool默认 True为 True 时忽略 run_number改用int(time.time() * 1e7) % 2**31随机取种--load_replay_dir字符串默认 None从固定数据集目录加载初始回放缓冲None 则不从外部加载--load_replay_numberint默认 None加载固定回放数据时使用的运行编号默认沿用run_number--save_replaybool默认 False训练结束后将最终回放缓冲保存为固定数据集到${base_dir}/replay_logs--data_loggingbool默认 False是否用智能体记录回放缓冲当前实现直接抛出NotImplementedError--max_episode_evalbool默认 True是否使用按固定 episode 数评估的DataEfficientAtariRunner--tag字符串默认 None本次运行的标签会写入 config.json3.2 入口执行流程源码视角从 bbf/train.py 的main()可以看出完整启动链路设置 TensorFlow 行为、GPU 内存增长或隐藏 GPU当非run_xm_preprocessing路径时确定随机种子并调用set_random_seed()同时设置PYTHONHASHSEED、tf.random、np.randomrun_experiment.load_gin_configs()解析 gin 文件与--gin_bindings覆盖项write_config()将当前 gin 配置、seed、tag、agent 名落盘为base_dir/config.jsonbbf/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 参数继承自 DopamineJaxDQNAgentJaxDQNAgent.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 adamgamma 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 RatioBBFAgent.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-perturbBBFAgent.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_weightsbbf/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_schedulerbbf/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 5SPR 转移模型在潜空间向前预测 5 步回放缓冲以subseq_len jumps 1的子序列形式采样见 bbf/agents/spr_agent.pyspr_weight 5SPR 辅助损失权重总损失为loss dqn_loss spr_weight * spr_loss其中spr_loss ||spr_predictions - spr_targets||²并按轨迹掩码取平均bbf/agents/spr_agent.pydata_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优化器参数沿用 DERvan Hasselt 等2019Adam 的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_namePongsticky_actions FalseAtari 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 \ --agentSPR \ --gin_filesbbf/configs/SPR.gin \ --base_dir/tmp/online_rl/spr \ --run_number1通过对比 bbf/configs/SPR.gin 与 bbf/configs/SR_SPR.gin 可以直观看到算法即配置的哲学维度SPRSR-SPR回放比率replay_ratio64256更高重置间隔reset_every未启用5_000shrink_factor/perturb_factor—0.8 / 0.2batches_to_group未设置默认 18编码器dqnhidden_dim512width_scale1cnnhidden_dim512width_scale1噪声noisyTruenoisyFalse注释说明 noisy 更慢且损害性能权重衰减未设置默认 0.1 经注释注明参数源自 DER0.0更新视界10固定未显式设置依赖 reset 周期退火机制默认游戏BreakoutChopperCommandSPR 对应 2021 年的自预测表示论文SR-SPR 则是在其基础上加入高回放比率与周期性重置二者性能与行为差异完全由 gin 参数体现无需改动任何 Python 代码。六、评估机制DataEfficientAtariRunner 与归一化分数默认--max_episode_evalTrue时使用 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 summaryTrain/EpisodeReturn、Eval/NormalizedScore等。七、实验结果与消融数据scores 目录仓库 bigger_better_faster/scores 目录存放了 BBF 论文消融实验的原始结果 CSV按回放比率分组RR2 / RR8 主结果RR2_BBF.csv、RR8_BBF.csv消融变体文件命名即实验标签RR2_BBFsticky_20k.csv~RR2_BBFsticky_1M.csv在训练 2 万步至 100 万步区间启用粘性动作的对比RR2_BBFγ0.99.csv固定折扣因子 0.99对应去除 gamma 退火RR2_BBFn10.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_filesbbf/configs/BBF.gin是相对路径需在仓库根目录含bbf/包的目录即bigger_better_faster/下执行环境适配JAX 的 GPU/CUDA 安装是独立步骤务必按官方指引完成后再pip install -r requirements.txtgym版本被严格锁定为0.25.2且 Atari ROM 需通过accept-rom-license授权随机种子--no_seeding默认为 True随机取种需要可复现实验时显式传入--no_seedingFalse --run_numberN或--agent_seedN修改回放比率时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 DOro, 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-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载相关推荐Dopamine中的元学习快速适应新环境的RL算法Dopamine中的元学习快速适应新环境的RL算法 引言 在强化学习Reinforcement Learning, RL领域智能体Agent通常需要强化学习机器学习深度学习DreamerV2基于离散世界模型的强化学习框架技术深度解析DreamerV2基于离散世界模型的强化学习框架技术深度解析 在强化学习领域基于模型的强化学习框架正逐渐成为研究热点。DreamerV2作为这一领域的代表性AReaL TIR 智能体实战基于多轮工具调用的数学推理强化学习指南AReaL TIR 智能体实战基于多轮工具调用的数学推理强化学习指南 导读 本文聚焦 AReaL 仓库中 examples/tir 提供的 Tool Int人工智能大模型强化学习分布式训练AI Agent上一篇洛雪音乐音源架构深度解析构建高可用全网音乐聚合平台的技术实现下一篇autojump数据库分布式架构故障恢复流程与测试创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表