Dopamine 连续控制域训练 Runner:ContinuousTrainRunner 源码解析与实战指南
机器学习深度学习【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址https://gitcode.com/gh_mirrors/do/dopamine点击查看免费下载导读ContinuousTrainRunner是 Dopamine 强化学习框架中专门服务于连续动作空间continuous control实验的 Runner 类其核心使命是在不执行评估阶段的前提下驱动 JAX/Flax 智能体在 Gym 环境如 MuJoCo 的 HalfCheetah、Ant 等中持续训练。本文以该类的 API 文档为骨架结合仓库源码逐层剖析其继承结构、训练迭代流程、指标记录与 TensorBoard 汇总逻辑并给出基于 gin 配置文件和命令行入口的完整实操方案帮助你快速搭建只训练、不评估的连续域实验。类的定位专为 JAX/Flax 智能体设计的纯训练 Runner根据 API 文档 的定义ContinuousTrainRunner是 Object that handles running experiments负责运行实验的对象其关键说明是This is mostly the same as discrete_domains.TrainRunner, but is written solely for JAX/Flax agents.也就是说它与离散动作域中的TrainRunner在行为上基本一致——都取消了评估阶段、只保留训练循环——但ContinuousTrainRunner专门面向 JAX/Flax 智能体实现不涉及 TensorFlow Session 机制。从源码 run_experiment.py 可以看到其完整类定义与继承链gin.configurable class ContinuousTrainRunner(ContinuousRunner): Object that handles running experiments. This is mostly the same as discrete_domains.TrainRunner, but is written solely for JAX/Flax agents. 完整的继承层次如下dopamine.discrete_domains.run_experiment.Runner基类定义实验主循环、checkpoint、日志等通用逻辑见 run_experiment.pydopamine.continuous_domains.run_experiment.ContinuousRunner连续域 Runner见 API 文档 与 源码dopamine.continuous_domains.run_experiment.ContinuousTrainRunner本次主题纯训练模式其中ContinuousRunner同样被文档明确标注为 This is mostly the same as discrete_domains.Runner, but is written solely for JAX/Flax agents它承接了基类Runner的全部实验管理能力并针对连续域环境做了适配。构造函数三个参数与初始化动作ContinuousTrainRunner的构造函数在 源码 L272-L289 中定义gin.configurable def __init__( self, base_dir, create_agent_fn, create_environment_fngym_lib.create_gym_environment, ): Initialize the TrainRunner object in charge of running a full experiment. logging.info(Creating ContinuousTrainRunner ...) super().__init__(base_dir, create_agent_fn, create_environment_fn) self._agent.eval_mode False参数说明如下参数类型默认值含义base_dirstr无必填存放所有子目录checkpoints、logs、TensorBoard 事件文件的基础目录create_agent_fn函数无必填接收环境对象并返回一个智能体的工厂函数create_environment_fn函数gym_lib.create_gym_environment创建 Gym 环境的工厂函数构造函数做了三件关键事情打印Creating ContinuousTrainRunner ...日志调用父类ContinuousRunner的__init__完成目录创建、TensorBoard SummaryWriter 初始化、环境创建、智能体创建、checkpoint 恢复以及CollectorDispatcher指标分发器的挂载见 L187-L213显式将self._agent.eval_mode False确保智能体始终处于训练模式——这与TrainRunner在离散域中的做法完全一致见 discrete_domains/run_experiment.py L770。需要特别注意的是父类ContinuousRunner的__init__中create_agent_fn的调用签名是create_agent_fn(self._environment, summary_writerself._summary_writer)L200-L202因此在自定义create_agent_fn时应当兼容环境对象 关键字参数summary_writer的调用形式。迭代循环只跑训练阶段不做评估ContinuousTrainRunner相对父类最核心的差异体现在_run_one_iteration方法源码 L291-L328def _run_one_iteration(self, iteration): Runs one iteration of agent/environment interaction. An iteration involves running several episodes until a certain number of steps are obtained. This method differs from the _run_one_iteration method in the base Runner class in that it only runs the train phase. statistics iteration_statistics.IterationStatistics() num_episodes_train, average_reward_train, average_steps_per_second ( self._run_train_phase(statistics) ) if self._has_collector_dispatcher: self._collector_dispatcher.write([ statistics_instance.StatisticsInstance( Train/NumEpisodes, num_episodes_train, iteration ), statistics_instance.StatisticsInstance( Train/AverageReturns, average_reward_train, iteration ), statistics_instance.StatisticsInstance( Train/AverageStepsPerSecond, average_steps_per_second, iteration ), ]) self._save_tensorboard_summaries( iteration, num_episodes_train, average_reward_train, average_steps_per_second, ) return statistics.data_lists对比基类Runner._run_one_iterationdiscrete_domains/run_experiment.py L572-L620差异一目了然基类 Runner先_run_train_phase再_run_eval_phase收集训练与评估两套指标含Eval/NumEpisodes、Eval/AverageReturnsContinuousTrainRunner只调用_run_train_phase指标仅包含训练维度同时把iteration作为 TensorBoard 的 global step 使用。训练阶段内部做了什么_run_train_phase由基类Runner实现L515-L546其流程为将self._agent.eval_mode置为False记录开始时间调用_run_one_phase(self._training_steps, statistics, train)通过循环执行完整 episode直到累计步数达到training_steps计算平均未折扣回报average_return sum_returns / num_episodes与每秒训练步数average_steps_per_second number_steps / time_delta将train_average_return、train_average_steps_per_second追加进统计对象。每个 episode 的执行_run_one_episodeL390-L428遵循 Machado et al., 2017 的惯例跑完整 episode 直到终止若启用了奖励裁剪则在[-1, 1]区间内做np.clip。基类还支持max_steps_per_episodeNone时跨迭代续跑 episode 的_run_continued_episode模式L430-L468。指标记录CollectorDispatcher 与 TensorBoardContinuousTrainRunner覆盖了_save_tensorboard_summaries方法L330-L341仅向 TensorBoard 写入三个训练指标def _save_tensorboard_summaries( self, iteration, num_episodes, average_reward, average_steps_per_second ): Save statistics as tensorboard summaries. metrics [ (Train/NumEpisodes, num_episodes), (Train/AverageReturns, average_reward), (Train/AverageStepsPerSecond, average_steps_per_second), ] for name, value in metrics: self._summary_writer.scalar(name, value, iteration) self._summary_writer.flush()同时在_run_one_iteration中通过CollectorDispatcher见 dopamine/metrics/collector_dispatcher.py以StatisticsInstance形式上报相同的三个指标实现实验指标的收集与分发。这套指标体系与离散域TrainRunnerL791-L808完全对齐。指标汇总如下指标名称含义写入目标Train/NumEpisodes本迭代执行的训练 episode 数TensorBoard CollectorDispatcherTrain/AverageReturns本迭代平均未折扣训练回报TensorBoard CollectorDispatcherTrain/AverageStepsPerSecond训练吞吐量每秒步数TensorBoard CollectorDispatcher值得注意的是父类ContinuousRunner的_save_tensorboard_summariesL233-L261还会额外写入Eval/NumEpisodes与Eval/AverageReturns这正是训练评估Runner 与纯训练Runner 在指标输出上的分水岭。实验主循环与 checkpoint 行为ContinuousTrainRunner本身没有重写run_experiment因此复用基类Runner.run_experiment的逻辑L714-L739若num_iterations start_iteration则直接返回否则用tqdm进度条迭代range(start_iteration, num_iterations)逐次调用_run_one_iteration每次迭代后旧版 Logger 记录实验数据_log_experiment、执行 checkpoint_checkpoint_experiment、刷新CollectorDispatcher实验结束前 flush TensorBoard 并关闭CollectorDispatcher。checkpoint 恢复机制由_initialize_checkpointer_and_maybe_resumeL307-L353实现它会检查base_dir/checkpoints下是否存在已完成的迭代编号若存在则调用智能体的unbundle恢复网络权重并从current_iteration 1继续训练。这意味着中断的训练可以直接续跑无需重新开始。工厂函数与 schedule 选择ContinuousTrainRunner通常不直接实例化而是通过create_continuous_runner工厂函数创建API 文档 与 源码 L111-L133gin.configurable def create_continuous_runner(base_dir, schedulecontinuous_train_and_eval): Creates an experiment Runner. assert base_dir is not None # Continuously runs training and evaluation until max num_iterations is hit. if schedule continuous_train_and_eval: return ContinuousRunner(base_dir, create_continuous_agent) # Continuously runs training until max num_iterations is hit. elif schedule continuous_train: return ContinuousTrainRunner(base_dir, create_continuous_agent) else: raise ValueError(Unknown schedule: {}.format(schedule))schedule取值返回对象行为continuous_train_and_eval默认ContinuousRunner每迭代先训练再评估continuous_trainContinuousTrainRunner每迭代只训练不评估其他值抛出ValueError报错Unknown schedule: ...智能体工厂create_continuous_agent与ContinuousTrainRunner配套的智能体工厂是create_continuous_agentAPI 文档 与 源码 L42-L108它根据agent_name分发到对应实现agent_name返回的智能体前置条件sacSACAgent动作与观测空间均为spaces.box.BoxppoPPOAgent动作与观测空间均为spaces.Boxsac_cale*前缀匹配SACCALEAgent来自 dopamine/labs/cale/sac_cale.pyppo_cale*前缀匹配PPOCALEAgent来自 dopamine/labs/cale/ppo_cale.py其他抛出ValueError: Unknown agent: ...—其中SACAgent的构造需要action_shape、action_limits由action_space.low/high提供、observation_shape、动作/观测的 dtype 以及可选的summary_writer供智能体内部把训练统计写入 TensorBoard。命令行入口train.py 与 gin 配置实战命令行入口解析连续域实验的官方入口是 dopamine/continuous_domains/train.py它定义并解析三个命令行参数flags.DEFINE_string( base_dir, None, Base directory to host all required sub-directories. ) flags.DEFINE_multi_string( gin_files, [], List of paths to gin configuration files (e.g. dopamine/jax/agents/sac/configs/sac.gin)., ) flags.DEFINE_multi_string( gin_bindings, [], Gin bindings to override the values set in the config files., )main函数的执行链路为run_experiment.load_gin_configs(gin_files, gin_bindings) runner run_experiment.create_continuous_runner(base_dir) runner.run_experiment()即加载 gin 配置文件 → 用create_continuous_runnerschedule 由 gin 绑定决定创建 Runner → 调用run_experiment启动训练。注意base_dir被标记为必填flags.mark_flag_as_required(base_dir)未指定会直接报错。一个典型的启动命令形如python -m dopamine.continuous_domains.train \ --base_dir/tmp/dopamine/continuous/sac_halfcheetah \ --gin_filesdopamine/jax/agents/sac/configs/sac.gingin 配置实战以 SAC 为例配置文件 dopamine/jax/agents/sac/configs/sac.gin 展示了与ContinuousTrainRunner直接相关的关键绑定create_gym_environment.environment_name HalfCheetah create_gym_environment.version v2 create_continuous_runner.schedule continuous_train_and_eval create_continuous_agent.agent_name sac ContinuousTrainRunner.create_environment_fn gym_lib.create_gym_environment ContinuousRunner.num_iterations 3_200 ContinuousRunner.training_steps 1_000 ContinuousRunner.evaluation_steps 10_000 # agent steps ContinuousRunner.max_steps_per_episode 1_000 ContinuousRunner.clip_rewards False要点解读环境通过create_gym_environment绑定到 MuJoCo 的HalfCheetah-v2。create_gym_environmentgym_lib.py L58-L100支持的环境列表见MUJOCO_GAMES (Ant, HalfCheetah, Hopper, Humanoid, Walker2d)L55并会剥离 Gym 的TimeLimit包装器默认将 episode 限制在 200 步随后包一层GymPreprocessingL103-L143以符合 Dopamine 的 API 约定schedule 绑定gin 绑定create_continuous_runner.schedule默认continuous_train_and_eval返回ContinuousRunner若改为continuous_train则返回ContinuousTrainRunner从此跳过评估阶段ContinuousTrainRunner.create_environment_fn显式绑定环境工厂确保两种 Runner 都使用同一环境创建函数训练规模num_iterations3_200、training_steps1_000每迭代训练步数、max_steps_per_episode1_000、clip_rewardsFalse连续域通常不裁剪奖励。若希望把上述 SAC 实验切换为纯训练模式只需增加一行 gin 绑定即可无需改动任何 Python 代码python -m dopamine.continuous_domains.train \ --base_dir/tmp/dopamine/continuous/sac_halfcheetah_train_only \ --gin_filesdopamine/jax/agents/sac/configs/sac.gin \ --gin_bindingscreate_continuous_runner.schedulecontinuous_train测试验证行为契约的可信依据仓库中的单元测试 tests/dopamine/continuous_domains/run_experiment_test.py 为上述行为提供了直接验证testCreateContinuousAgentReturnsAgentcreate_continuous_agent(env, sac)返回SACAgent实例testCreateContinuousAgentWithInvalidNameRaisesException非法智能体名抛出ValueErrortestCreateContinuousRunnerCreatesCorrectRunner参数化schedulecontinuous_train_and_eval返回ContinuousRunnerschedulecontinuous_train返回ContinuousTrainRunnertestCreateContinuousRunnerFailsWithInvalidName非法 schedule 抛出ValueError。测试中通过gin.bind_parameter(ContinuousTrainRunner.create_environment_fn, lambda: self.env)注入 mock 环境印证了ContinuousTrainRunner是 gin-configurable 的、且环境工厂可以自由替换。何时选择 ContinuousTrainRunner综合以上源码事实ContinuousTrainRunner适合以下场景训练为主、评估另行处理当评估开销较大如 MuJoCo 连续域评估需要大量 episode或评估由外部脚本统一执行时用它跑纯训练可以最大化训练吞吐调试与快速迭代需要快速验证算法稳定性、复现曲线时跳过评估阶段能显著缩短单次迭代时间与离散域 TrainRunner 对称的连续域方案Dopamine 在 discrete_domains/run_experiment.py 中为离散域提供了同构的TrainRunner连续域则由ContinuousTrainRunner补齐对应能力二者在schedule命名continuous_train与行为语义上保持一致。需要完整训练评估交替实验时应回退到默认的ContinuousRunnerschedule 保持continuous_train_and_eval需要跨迭代评估曲线用于论文复现时同样建议使用默认 Runner。选择的关键判断标准只有一个当前实验是否需要每轮迭代都产出评估指标。参考资源索引类文档ContinuousTrainRunner.md父类文档ContinuousRunner.md工厂函数文档create_continuous_runner.md、create_continuous_agent.md核心实现dopamine/continuous_domains/run_experiment.py、dopamine/continuous_domains/train.py基类实现dopamine/discrete_domains/run_experiment.py环境工厂dopamine/discrete_domains/gym_lib.py配置示例dopamine/jax/agents/sac/configs/sac.gin、dopamine/jax/agents/ppo/configs/ppo.gin测试用例tests/dopamine/continuous_domains/run_experiment_test.py赞分享机器学习深度学习【免费下载链接】dopamineDopamine is a research framework for fast prototyping of reinforcement learning algorithms.项目地址https://gitcode.com/gh_mirrors/do/dopamine点击查看免费下载相关推荐Dopamine 连续控制域实验入口create_continuous_runner 完整解析与实战指南Dopamine 连续控制域实验入口create_continuous_runner 完整解析与实战指南 导读 dopamine.continuous_dom机器学习深度学习Dopamine 连续控制域实验运行器 ContinuousRunner 完全指南JAX/Flax Agent 的训练调度、参数配置与源码剖析Dopamine 连续控制域实验运行器 ContinuousRunner 完全指南JAX/Flax Agent 的训练调度、参数配置与源码剖析 导读 dopa机器学习深度学习Dopamine create_continuous_agent 深度解析连续控制域 RL Agent 的统一工厂与 Gin 配置实战Dopamine create_continuous_agent 深度解析连续控制域 RL Agent 的统一工厂与 Gin 配置实战 导读 在 Dopami机器学习深度学习创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
上一篇/下一篇内容由系统自动关联
返回资讯列表 →