ARTICLE DETAIL

资讯详情

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

Dopamine 实验运行框架解析:run_experiment 模块的 Runner、TrainRunner 与完整训练流程

Dopamine 实验运行框架解析:run_experiment 模块的 Runner、TrainRunner 与完整训练流程 Dopamine 实验运行框架解析run_experiment 模块的 Runner、TrainRunner 与完整训练流程【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址: https://gitcode.com/gh_mirrors/dopami/dopaminedopamine.discrete_domains.run_experiment是 Dopamine 强化学习研究框架中负责运行实验的核心模块它定义了通用 Agent 的辅助类与方法把环境创建、Agent 实例化、训练/评估循环、数据记录与断点续训封装成一套标准流程。本文以 run_experiment.md 的 API 说明为主线结合 run_experiment.py 的源码实现、gin 配置文件与集成测试完整讲解如何用一行命令启动 DQN/Rainbow/Implicit Quantile 等 Agent 的 Atari 实验以及如何通过配置参数掌控整个实验生命周期。读完本文你将掌握 Runner 的内部运行机制、训练/评估阶段的调度逻辑、断点恢复原理并能独立编写自己的实验启动脚本。模块定位实验的总指挥在 Dopamine 中实验experiment一词特指模拟 Agent 与环境之间的交互并报告这些交互产生的统计指标这一完整过程见 run_experiment.py 的类 docstring。run_experiment模块处于离散域discrete domains实验栈的最上层向下调度以下组件环境创建默认使用 atari_lib.create_atari_environment 创建 Gym 风格的 Atari 2600 环境Agent 实例化通过create_agent工厂按名称装配框架内置的多种 Agent断点管理通过 checkpointer.Checkpointer 保存/恢复实验状态数据记录通过 iteration_statistics.IterationStatistics 收集指标经旧版 Logger 或新的 CollectorDispatcher 落盘。模块对外暴露两个类Runner、TrainRunner和两个工厂函数create_agent、create_runner全部由 gin 配置驱动。类与函数总览根据模块文档 run_experiment.md该模块的公开 API 结构如下名称类型作用Runnerclass负责运行 Dopamine 实验含训练 评估两个阶段TrainRunnerclass只运行训练阶段、不做评估的实验 Runnercreate_agent(...)function创建 RL Agentcreate_runner(...)function创建实验 Runner其中TrainRunner继承自Runner仅覆盖迭代逻辑TrainRunner与基类Runner的区别在于它不执行评估阶段但训练阶段的检查点保存与日志记录行为完全保留见 TrainRunner.md。用代码跑通第一个实验Runner 的最小使用示例Runner.md 给出了训练一个 DQN Agent 的最小示例import dopamine.discrete_domains.atari_lib base_dir /tmp/simple_example def create_agent(sess, environment): return dqn_agent.DQNAgent(sess, num_actionsenvironment.action_space.n) runner Runner(base_dir, create_agent, atari_lib.create_atari_environment) runner.run()这段代码演示了 Runner 的三个关键约定base_dir实验所有子目录checkpoints、logs的宿主根目录create_agent_fn一个以 TensorFlow Session 和环境为参数、返回 Agent 的函数——它解耦了如何构造 Agent与如何运行实验create_environment_fn默认即atari_lib.create_atari_environment接收问题名并创建对应 Gym 环境。需要说明的是在 run_experiment.py 的当前实现中Runner构造时并不会立刻创建 Session——注释明确说明Agent 现在负责设置 Sessionself._sess None构造 Agent 时传入summary_writerself._base_dir以目录代替真正的 SummaryWriter由 Agent 内部自行创建随后从self._agent._sess取回 Session见 run_experiment.py。这与上述示例中先建 sess 再建 agent的旧式写法略有出入实际运行时以 run_experiment.py 源码为准。Runner 构造参数全解Runner.__init__的参数定义见 run_experiment.py构成了实验的完整控制面。结合 dqn.gin 中的默认绑定各参数含义如下参数默认值说明base_dir必填承载所有子目录的基目录create_agent_fn必填接收 TensorFlow Session 与环境、返回 Agent 的函数create_environment_fnatari_lib.create_atari_environment接收问题名并创建 Gym 环境的函数如 Atari 2600 游戏checkpoint_file_prefixckpt检查点文件前缀logging_file_prefixlog日志文件前缀log_every_n1写日志的频率每 N 个迭代写一次num_iterations200迭代次数阈值必须大于start_iterationtraining_steps250000训练步数evaluation_steps125000评估步数max_steps_per_episode27000单个 episode 达到该步数后强制终止clip_rewardsTrue是否将奖励裁剪到 [-1, 1]use_legacy_loggerTrue是否使用旧版 Logger即将被 CollectorDispatcher 取代fine_grained_print_to_consoleTrue是否向控制台打印细粒度进度便于调试构造函数会依次执行以下初始化动作源码注释明确列出初始化一个环境初始化tf.compat.v1.Session由 Agent 内部创建并取回初始化 logger初始化 agent若存在最新检查点则恢复并初始化 Checkpointer 对象见 run_experiment.py。此外Runner 还会创建CollectorDispatcher用于新的指标上报体系并通过set_collector_dispatcher回调注入 Agent见 run_experiment.py。一次迭代的完整生命周期Runner 的核心运行单元是迭代iteration。_run_one_iteration见 run_experiment.py将一个迭代拆为训练与评估两个阶段其节奏设计用于对齐 Nature DQNMnih et al., 2015的 train/eval 交错方式。阶段一训练阶段_run_train_phase执行流程见 run_experiment.py将agent.eval_mode置为False调用_run_one_phase(self._training_steps, statistics, train)持续跑完整 episode直到累计步数达到training_steps计算并记录train_average_return平均未折扣回报与train_average_steps_per_second每秒训练步数。阶段二评估阶段_run_eval_phase执行流程见 run_experiment.py将agent.eval_mode置为True不学习运行_run_one_phase(self._evaluation_steps, statistics, eval)记录eval_average_return。底层循环_run_one_episode 与 _run_one_phase单 episode 的执行遵循 Machado et al., 2017 的约定跑完整 episode累积步数达到最小步数阈值才结束见 run_experiment.py。_run_one_episode见 run_experiment.py的细节值得注意用_initialize_episode()拿到初始动作进入交互循环每步返回observation, reward, is_terminal累加total_reward与step_number若clip_rewardsTrue则reward np.clip(reward, -1, 1)终止条件environment.game_over为真或步数达到max_steps_per_episode时结束 episode丢命处理若is_terminal为真但 episode 未结束Atari 中失去一条命则向 Agent 发送人工 episode 结束信号_end_episode(reward, is_terminal)然后调用begin_episode开启新一命继续否则正常调用agent.step(reward, observation)。注意_end_episode对 JAX Agent 与 TF Agent 的差异化处理JaxDQNAgent 支持传入terminal信号而 TF 系 Agent 仅传入 reward见 run_experiment.py。_run_one_phase则在每跑完一个 episode 后把train_episode_lengths/train_episode_returns或eval_前缀追加进IterationStatistics并在fine_grained_print_to_console开启时用sys.stdout.write实时刷新进度步数、episode 长度、回报。迭代收尾指标上报、日志与检查点一个迭代结束后见_run_one_iteration与run_experimentrun_experiment.py通过CollectorDispatcher.write上报 5 项标量Train/NumEpisodes、Train/AverageReturns、Train/AverageStepsPerSecond、Eval/NumEpisodes、Eval/AverageReturns若存在 SummaryWriter则写入 TensorBoard summariesTF1 分支用tf.compat.v1.Summary否则用tf.summary.scalar旧版 Logger 每log_every_n个迭代调用log_to_file(logging_file_prefix, iteration)调用_checkpoint_experiment(iteration)保存检查点将agent.bundle_and_checkpoint()的产物连同current_iteration、logs一起交给checkpointer.save_checkpoint()见 run_experiment.py。TrainRunner跳过评估、专注训练的变体TrainRunner见 run_experiment.py继承自Runner其差异化行为有两处构造后立即将agent.eval_mode置为False保证全程处于训练模式重写_run_one_iteration只执行_run_train_phase不再调用评估阶段相应地仅上报 3 项Train/*指标。这对于只需要快速迭代模型、暂不需要周期性评估的场景非常实用例如在集成测试中加速跑通全流程。工厂函数create_agent 与 create_runnercreate_agent按名称装配 Agentcreate_agent(sess, environment, agent_nameNone, summary_writerNone, debug_modeFalse)见 run_experiment.py是一个gin.configurable函数按agent_name字符串分发到对应 Agent 实现agent_name实例化的 Agent后端dqndqn_agent.DQNAgentTensorFlowrainbowrainbow_agent.RainbowAgentTensorFlowimplicit_quantileimplicit_quantile_agent.ImplicitQuantileAgentTensorFlowjax_dqnjax_dqn_agent.JaxDQNAgentJAXjax_quantilejax_quantile_agent.JaxQuantileAgentJAXjax_rainbowjax_rainbow_agent.JaxRainbowAgentJAXfull_rainbowfull_rainbow_agent.JaxFullRainbowAgentJAXjax_implicit_quantilejax_implicit_quantile_agent.JaxImplicitQuantileAgentJAX其他抛出ValueError(Unknown agent: {})—参数要点sesstf.compat.v1.Session用于运行关联算子JAX 系 Agent 不使用environmentGym 环境Agent 依据environment.action_space.n确定动作数summary_writerTensorFlow summary writer用于 Agent 内训数据进 TensorBoarddebug_mode为True时输出 episode 级统计到 TensorBoard默认关闭因为会拖慢训练。源码中若debug_modeFalse会强制把summary_writer置为None。create_runner按调度策略创建 Runnercreate_runner(base_dir, schedulecontinuous_train_and_eval)见 run_experiment.py同样是gin.configurable函数支持两种调度continuous_train_and_eval默认返回Runner持续训练 评估直到max_num_iterationscontinuous_train返回TrainRunner持续训练直到max_num_iterations其他取值抛出ValueError(Unknown schedule: {})。用 gin 配置驱动整个实验从命令行到 Runnerrun_experiment模块本身不提供 CLI而是由入口脚本 train.py 承接。train.py定义三个命令行 flagpython -m dopamine.discrete_domains.train \ --base_dir/tmp/dopamine \ --gin_filesdopamine/agents/dqn/configs/dqn.gin \ --gin_bindingscreate_environment.game_namePong \ --gin_bindingsRunner.num_iterations200--base_dir必填所有子目录的宿主目录flags.mark_flag_as_required(base_dir)强制校验--gin_filesmulti_stringgin 配置文件路径列表--gin_bindingsmulti_string覆盖配置文件中参数的 gin 绑定例如DQNAgent.epsilon_train0.1、create_environment.game_namePong。main的调用链非常简洁见 train.pymain └─ tf.compat.v1.disable_v2_behavior() # 禁用 TF2 行为兼容 TF1 风格代码 └─ run_experiment.load_gin_configs(gin_files, gin_bindings) # 解析 gin 配置 └─ run_experiment.create_runner(base_dir) # 按 schedule 创建 Runner └─ Runner / TrainRunner(base_dir, create_agent) └─ runner.run_experiment() # 启动迭代循环load_gin_configs调用gin.parse_config_files_and_bindings(gin_files, bindingsgin_bindings, skip_unknownFalse)见 run_experiment.pyskip_unknownFalse意味着配置中引用了未知符号会直接报错保证配置严谨性。一份可直接落地的 dqn.gin 配置解读以经典 DQN 配置 dqn.gin 为例它展示了 gin 如何把模块参数绑定到 Runner 与 Agent# 导入所需模块注册可配置符号 import dopamine.discrete_domains.atari_lib import dopamine.discrete_domains.run_experiment import dopamine.agents.dqn.dqn_agent import dopamine.replay_memory.circular_replay_buffer import gin.tf.external_configurables DQNAgent.gamma 0.99 DQNAgent.update_horizon 1 DQNAgent.min_replay_history 20000 # agent steps DQNAgent.update_period 4 DQNAgent.target_update_period 8000 # agent steps DQNAgent.epsilon_train 0.01 DQNAgent.epsilon_eval 0.001 DQNAgent.epsilon_decay_period 250000 # agent steps DQNAgent.tf_device /gpu:0 # use /cpu:* for non-GPU version DQNAgent.optimizer tf.train.RMSPropOptimizer() tf.train.RMSPropOptimizer.learning_rate 0.00025 tf.train.RMSPropOptimizer.decay 0.95 tf.train.RMSPropOptimizer.momentum 0.0 tf.train.RMSPropOptimizer.epsilon 0.00001 tf.train.RMSPropOptimizer.centered True atari_lib.create_atari_environment.game_name Pong # Sticky actions with probability 0.25, as suggested by (Machado et al., 2017). atari_lib.create_atari_environment.sticky_actions True create_agent.agent_name dqn Runner.num_iterations 200 Runner.training_steps 250000 # agent steps Runner.evaluation_steps 125000 # agent steps Runner.max_steps_per_episode 27000 # agent steps WrappedReplayBuffer.replay_capacity 1000000 WrappedReplayBuffer.batch_size 32注意其中既有Runner.*这类直接作用于 Runner 构造参数的绑定也有create_agent.agent_name dqn这类作用于工厂函数的绑定还有WrappedReplayBuffer.*这类深入到回放缓冲区实现的绑定——这正是 gin 配置一处声明、处处生效的设计create_agent与Runner都标注了gin.configurable因此它们的所有构造参数都能被 gin 注入。类似地其他 Agent 的配置如 rainbow.gin、implicit_quantile.gin以及 JAX 系配置如 dopamine/jax/agents/dqn/configs/dqn.gin均遵循同一套绑定模式。断点续训与数据落盘机制检查点从任意迭代恢复Runner 的_initialize_checkpointer_and_maybe_resume见 run_experiment.py实现了自动恢复创建Checkpointer(checkpoint_dir, checkpoint_file_prefix)调用checkpointer.get_latest_checkpoint_number()查询最大检查点编号——检查点 0 的存在意味着迭代 0 已完成因此从迭代 1 开始若存在检查点加载数据并调用agent.unbundle(...)恢复 Agent 内部状态校验数据包含logs与current_iteration键将start_iteration设为current_iteration 1。检查点的落盘细节可参考 checkpointer.py 的文档每次迭代写入ckpt.#文件并额外写入sentinel_checkpoint_complete.#哨兵文件标记全局保存成功因此 Agent 必须先保存图与回放缓冲区、后调用save_checkpointCheckpointer 还会清理旧文件仅保留最近CHECKPOINT_DURATION个迭代。统计指标IterationStatistics Logger CollectorDispatcheriteration_statistics.py 中的IterationStatistics维护data_lists键到列表的映射例如train_episode_returns记录本迭代各 episode 的回报旧版 logger.py 的Logger用logging_dir存储数据字典按log_{iteration}文件落盘默认保留logs_duration4份Runner 侧默认use_legacy_loggerTrue且_create_directories中会打印弃用警告提示切换到CollectorDispatcher新体系下Runner 每迭代通过CollectorDispatcher.write提交StatisticsInstance标量并定期flush实验结束时close。集成测试验证 Runner 全流程的证据仓库中的 train_runner_integration_test.py 从端到端角度验证了 Runner 的约定可作为上文的实证通过FLAGS.gin_files [dopamine/jax/agents/dqn/configs/dqn.gin]指定 JAX DQN 配置用gin_bindings将实验缩到极小规模create_runner.schedulecontinuous_train使用 TrainRunner、Runner.training_steps100、Runner.num_iterations1、Runner.max_steps_per_episode100、WrappedReplayBuffer.replay_capacity100等调用train.main([])跑完整入口断言checkpoints/ckpt.0、checkpoints/sentinel_checkpoint_complete.0、logs/log_0三个文件确实生成——分别印证了检查点文件、哨兵文件与日志文件的落盘行为。该测试同时说明任何自定义 Agent 只要遵循create_agent_fn的签名约定就能无缝接入Runner的实验框架。总结理解 run_experiment 的要点实验 迭代循环每次迭代内含跑满training_steps的训练阶段 跑满evaluation_steps的评估阶段直到num_iterations或断点恢复后的起点两个 RunnerRunner训练 评估与TrainRunner仅训练由create_runner的schedule参数选择一切皆可 gin 配置create_agent、create_runner、Runner、Agent 与回放缓冲区均通过gin.configurable暴露参数train.py 的--gin_files与--gin_bindings是唯一入口断点续训开箱即用ckpt.# 哨兵文件保证原子性current_iteration决定恢复起点指标双轨制旧版Logger与新版CollectorDispatcher并存前者即将弃用。若需更细粒度参考可继续阅读 Runner 类 API、create_agent API 与 create_runner API并对照 run_experiment.py 源码逐行研读。【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址: https://gitcode.com/gh_mirrors/dopami/dopamine创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表