恒美微站
首页
关于我们
建站服务
主题模板
案例展示
资讯中心
联系我们
Tianshou 离线强化学习完全指南:D4RL 连续控制与 Atari 离散控制的离线训练实战
首页
资讯中心
/
Tianshou 离线强化学习完全指南:D4RL 连续控制与 Atari 离散控制的离线训练实战
Tianshou 离线强化学习完全指南:D4RL 连续控制与 Atari 离散控制的离线训练实战
发布时间:2026/9/23 21:47:08
人工智能机器学习深度学习强化学习【免费下载链接】tianshouAn elegant PyTorch deep reinforcement learning library.项目地址https://gitcode.com/gh_mirrors/ti/tianshou点击查看免费下载离线强化学习Offline Reinforcement Learning设定中智能体从一份一次性采集完成、此后不再与环境交互的固定数据集fixed dataset中学习策略。Tianshou 在 examples/offline 目录下提供了完整的离线 RL 示例套件连续控制方向基于 D4RL 数据集实现了 BCQ、CQL、TD3BC 与行为克隆IL四类算法离散控制方向则基于 Atari 专家数据实现了 IL、BCQ、CQL、CRR 四类算法并附带 RL Unplugged 数据集的转换脚本。读完本文你将掌握offline_trainer的用法、D4RL 数据的加载方式、Atari 专家数据的采集流程、各算法关键超参数的调优建议以及如何复现本文档中记录的全部基准结果。离线 RL 的基本设定与 Tianshou 的支持与在线 RL智能体一边探索环境一边学习不同离线 RL 的训练数据是预先用任意策略采集好的一份固定数据集。训练一旦开始数据集不再变化智能体也不再与环境进行任何交互。这一设定使离线 RL 特别适合数据昂贵、环境不可用或安全要求高的真实场景。在 Tianshou 中离线训练的核心入口是offline_trainer在 tianshou/trainer.py 中实现。从源码看OfflineTrainer继承自Trainer基类其训练循环非常直观每个训练步从buffer中采样batch_size条转移并调用算法的update方法执行一次梯度更新整个过程完全不涉及环境交互# tianshou/trainer.py 中 OfflineTrainer._training_step 的核心逻辑 training_stats self.algorithm.update( sample_sizeself.params.batch_size, bufferself._buffer ) self._update_moving_avg_stats_and_log_update_data(training_stats) return self._TrainingStepResult( training_statstraining_stats, env_step_advancementself.params.batch_size, )OfflineTrainerParams的结构也很简单核心字段如下参数类型默认值含义bufferReplayBuffer必填用于离线训练的环境转移数据。训练前会经过算法的预处理函数如有batch_sizeint64每次梯度更新从 buffer 中采样的转移数量实际使用中还会配合max_epochs、epoch_num_steps、test_step_num_episodes、save_best_fn、logger等参数这些继承自TrainerParams共同构成一次完整的离线训练流程。由于sample_sizebatch_size每个训练步恰好执行一次梯度更新因此epoch_num_steps即代表每个 epoch 的梯度更新次数——这也是本文档中一个 epoch 表示 10k 梯度步说法的来源。连续控制D4RL 数据集上的离线训练对于连续控制任务Tianshou 使用 D4RL 数据集训练离线智能体。D4RL 是离线 RL 领域的标准基准数据集集合涵盖 Gym-MuJoCo、Adroit、AntMaze 等多个系列。使用 D4RL 前需要先按照其项目说明完成安装和数据集下载。Tianshou 为连续控制提供了BCQ和CQL两个算法的完整实现此外examples/offline目录中还额外提供了TD3BC与IL行为克隆示例d4rl_bcq.pyBCQ 算法示例d4rl_cql.pyCQL 算法示例d4rl_td3_bc.pyTD3BC 算法示例d4rl_il.py行为克隆IL示例将 D4RL 数据集解析为 ReplayBuffer所有 D4RL 示例的起点都是 examples/offline/utils.py 中的load_buffer_d4rl函数。它调用 D4RL 的qlearning_dataset接口把原始数据集字段observations、actions、rewards、terminals、next_observations映射为 Tianshou 的ReplayBufferdef load_buffer_d4rl(expert_data_task: str) - ReplayBuffer: dataset d4rl.qlearning_dataset(gym.make(expert_data_task)) return ReplayBuffer.from_data( obsdataset[observations], actdataset[actions], rewdataset[rewards], donedataset[terminals], obs_nextdataset[next_observations], terminateddataset[terminals], truncatednp.zeros(len(dataset[terminals])), )注意这里显式将truncated全部置零只保留terminals真终止作为 done 信号——这是 D4RL 数据与 Gymnasium 新 API 之间的适配要点。加载得到的ReplayBuffer将作为OfflineTrainerParams的buffer参数传入训练流程。BCQ 示例的运行与网络结构以 d4rl_bcq.py 为例训练命令为python3 d4rl_bcq.py --task HalfCheetah-v2 --expert-data-task halfcheetah-expert-v2脚本内部依次完成以下工作创建环境并读取空间信息通过gym.make(args.task)创建 MuJoCo 环境用SpaceInfo.from_env(env)读取观测形状、动作维度与动作范围再构建SubprocVectorEnv作为测试环境构建 BCQ 的三套网络对应 BCQ 原论文的架构Perturbation扰动网络net_aPerturbation输入拼接后的状态-动作输出带phi上界约束的扰动两个ContinuousCriticClipped Double Q-learning权重由lmbda控制VAE变分自编码器编码器输入状态-动作对解码器根据状态与隐变量重构动作latent_dim默认为2 * action_dim组装BCQPolicy与BCQ算法设置gamma、tau、lmbda、phi等超参数加载数据并启动离线训练调用load_buffer_d4rl得到replay_buffer传入OfflineTrainerParams执行algorithm.run_training(...)评估训练结束后在测试环境上采集num_test_envs个 episode打印统计结果。BCQ 脚本的关键超参数及其默认值参数默认值说明--taskHalfCheetah-v2MuJoCo 环境名--expert_data_taskhalfcheetah-expert-v2D4RL 数据集名--buffer_size1000000回放缓冲区容量--hidden_sizes[256, 256]网络隐藏层结构--actor_lr/--critic_lr1e-3actor / critic 学习率--epoch/--epoch_num_steps200/5000训练轮数与每轮梯度步数--n_step3N 步 TD 目标步数--batch_size256梯度更新批大小--vae_hidden_sizes[512, 512]VAE 编解码器隐藏层--latent_dim2 * action_dimVAE 隐变量维度--gamma/--tau0.99/0.005折扣因子 / 软更新系数--lmbda0.75Clipped Double Q-learning 中的权重系数--phi0.05BCQ 最大扰动超参数CQL 与 TD3BC 的关键差异d4rl_cql.py 的架构与 BCQ 明显不同它以SAC 风格的概率 Actor-Critic 为基础ContinuousActorProbabilistic 双ContinuousCritic外层套上CQL算法。CQL 特有的超参数包括--alpha默认0.2与--auto_alpha默认开启熵项权重AutoAlpha会将目标熵设为-action_dim进行自动调节--cql_weight默认1.0CQL 损失项的权重--with_lagrange默认True与--lagrange_threshold默认10.0是否使用拉格朗日乘子约束 Q 值以及乘子的阈值上限--calibrated默认TrueCQL 是否采用校准版本--cql_alpha_lr默认3e-4CQL 熵调节的学习率--temperature默认1.0Boltzmann 策略的温度系数。d4rl_td3_bc.py 则在TD3 基础上叠加行为克隆正则项ContinuousDeterministicPolicy携带GaussianNoise(sigma0.1)探索噪声TD3BC算法使用policy_noise0.2、noise_clip0.5、update_actor_freq2等 TD3 标准配置外加--alpha默认2.5作为行为克隆项权重。连续控制基准结果与 Observation Normalization文档记录的 HalfCheetah-v2 连续控制基准结果如下均使用--seed 0与默认超参数。IL行为克隆EnvironmentDatasetILParametersHalfCheetah-v2halfcheetah-expert-v211355.31python3 d4rl_il.py --task HalfCheetah-v2 --expert-data-task halfcheetah-expert-v2HalfCheetah-v2halfcheetah-medium-v25098.16python3 d4rl_il.py --task HalfCheetah-v2 --expert-data-task halfcheetah-medium-v2BCQEnvironmentDatasetBCQParametersHalfCheetah-v2halfcheetah-expert-v211509.95python3 d4rl_bcq.py --task HalfCheetah-v2 --expert-data-task halfcheetah-expert-v2HalfCheetah-v2halfcheetah-medium-v25147.43python3 d4rl_bcq.py --task HalfCheetah-v2 --expert-data-task halfcheetah-medium-v2CQLEnvironmentDatasetCQLParametersHalfCheetah-v2halfcheetah-expert-v22864.37python3 d4rl_cql.py --task HalfCheetah-v2 --expert-data-task halfcheetah-expert-v2HalfCheetah-v2halfcheetah-medium-v26505.41python3 d4rl_cql.py --task HalfCheetah-v2 --expert-data-task halfcheetah-medium-v2TD3BCEnvironmentDatasetTD3BCParametersHalfCheetah-v2halfcheetah-expert-v211788.25python3 d4rl_td3_bc.py --task HalfCheetah-v2 --expert-data-task halfcheetah-expert-v2HalfCheetah-v2halfcheetah-medium-v25741.13python3 d4rl_td3_bc.py --task HalfCheetah-v2 --expert-data-task halfcheetah-medium-v2Observation Normalization 的影响参照 TD3BC 原论文的做法TD3BC 示例默认开启观测归一化observation normalization可通过--norm-obs 0关闭。从实现看归一化包含两步见 examples/offline/utils.py 的normalize_all_obs_in_replay_buffer与 d4rl_td3_bc.py先用RunningMeanStd基于整个数据集统计观测均值与方差再把 buffer 中的obs/obs_next归一化为零均值单位方差同时把同一obs_rms设置到测试环境的VectorEnvNormObs包装器上update_obs_rmsFalse即测试时不再更新统计量。文档记录的开/关归一化对比结果如下差异较小但方向一致Datasetw/ norm-obsw/o norm-obshalfcheeta-medium-v25741.135724.41halfcheeta-expert-v211788.2511665.77walker2d-medium-v24051.763985.59walker2d-expert-v25068.155027.75离散控制基于 Atari 专家数据的离线训练对于离散控制Tianshou 目前使用由一个已训练的 QRDQN 智能体采集的 ad hoc Atari 数据进行训练支持 IL、BCQ、CQL、CRR 四种离线算法。数据采集三步流程要在 Atari 上运行 CQL 等离线算法需要依次完成以下三步训练专家智能体使用 Atari 示例中 QRDQN 一节的命令python3 atari_qrdqn.py --task {your_task}生成带噪声的专家 buffer用--watch模式回放预训练策略并采集数据python3 atari_qrdqn.py --task {your_task} --watch --resume-path log/{your_task}/qrdqn/policy.pth --eps-test 0.2 --buffer-size 1000000 --save-buffer-name expert.hdf5注意100 万条 Atari 转移无法以.pkl格式保存文件过大且会报错必须保存为.hdf5格式。训练离线模型python3 atari_{bcq,cql,crr}.py --task {your_task} --load-buffer-name expert.hdf5从 atari_bcq.py 等脚本的实现可以看到buffer 的加载逻辑同时支持两种格式.pkl文件直接用pickle读取.hdf5文件则通过VectorReplayBuffer.load_hdf5加载若文件不存在脚本会给出提示Please run atari_dqn.py first to get experts data buffer。网络方面Atari 图像观测N_FRAMES x H x W默认4帧堆叠先经过DQNet/QRDQNet卷积特征提取再接DiscreteActor或DiscreteCritic输出层。IL / BCQ / CQL / CRR 基准结果以下结果均基于v4 版本 Atari 环境与原作者使用的 v0 不同且一个 epoch 表示 10k 梯度步。表中同时给出在线 QRDQN 得分Online与行为策略Behavioral即专家数据本身的表现作为参照。ILTaskOnline QRDQNBehavioralILparametersPongNoFrameskip-v420.56.820.0 (epoch 5)python3 atari_il.py --task PongNoFrameskip-v4 --load-buffer-name log/PongNoFrameskip-v4/qrdqn/expert.hdf5 --epoch 5BreakoutNoFrameskip-v4394.346.9121.9 (epoch 12, could be higher)python3 atari_il.py --task BreakoutNoFrameskip-v4 --load-buffer-name log/BreakoutNoFrameskip-v4/qrdqn/expert.hdf5 --epoch 12BCQTaskOnline QRDQNBehavioralBCQparametersPongNoFrameskip-v420.56.820.1 (epoch 5)python3 atari_bcq.py --task PongNoFrameskip-v4 --load-buffer-name log/PongNoFrameskip-v4/qrdqn/expert.hdf5 --epoch 5BreakoutNoFrameskip-v4394.346.964.6 (epoch 12, could be higher)python3 atari_bcq.py --task BreakoutNoFrameskip-v4 --load-buffer-name log/BreakoutNoFrameskip-v4/qrdqn/expert.hdf5 --epoch 12CQLTaskOnline QRDQNBehavioralCQLparametersPongNoFrameskip-v420.56.820.4 (epoch 5)python3 atari_cql.py --task PongNoFrameskip-v4 --load-buffer-name log/PongNoFrameskip-v4/qrdqn/expert.hdf5 --epoch 5BreakoutNoFrameskip-v4394.346.9129.4 (epoch 12)python3 atari_cql.py --task BreakoutNoFrameskip-v4 --load-buffer-name log/BreakoutNoFrameskip-v4/qrdqn/expert.hdf5 --epoch 12 --min-q-weight 50缩减数据量实验将离线数据缩减到上述规模的 10% 与 1% 后重新训练min_q_weight需要随数据量调整Buffer size 100000TaskOnline QRDQNBehavioralCQLparametersPongNoFrameskip-v420.55.821 (epoch 5)python3 atari_cql.py --task PongNoFrameskip-v4 --load-buffer-name log/PongNoFrameskip-v4/qrdqn/expert.size_1e5.hdf5 --epoch 5BreakoutNoFrameskip-v4394.341.440.8 (epoch 12)python3 atari_cql.py --task BreakoutNoFrameskip-v4 --load-buffer-name log/BreakoutNoFrameskip-v4/qrdqn/expert.size_1e5.hdf5 --epoch 12 --min-q-weight 20Buffer size 10000TaskOnline QRDQNBehavioralCQLparametersPongNoFrameskip-v420.5nan1.8 (epoch 5)python3 atari_cql.py --task PongNoFrameskip-v4 --load-buffer-name log/PongNoFrameskip-v4/qrdqn/expert.size_1e4.hdf5 --epoch 5 --min-q-weight 1BreakoutNoFrameskip-v4394.331.722.5 (epoch 12)python3 atari_cql.py --task BreakoutNoFrameskip-v4 --load-buffer-name log/BreakoutNoFrameskip-v4/qrdqn/expert.size_1e4.hdf5 --epoch 12 --min-q-weight 10CRRTaskOnline QRDQNBehavioralCRRCRR w/ CQLparametersPongNoFrameskip-v420.56.8-21 (epoch 5)17.7 (epoch 5)python3 atari_crr.py --task PongNoFrameskip-v4 --load-buffer-name log/PongNoFrameskip-v4/qrdqn/expert.hdf5 --epoch 5BreakoutNoFrameskip-v4394.346.923.3 (epoch 12)76.9 (epoch 12)python3 atari_crr.py --task BreakoutNoFrameskip-v4 --load-buffer-name log/BreakoutNoFrameskip-v4/qrdqn/expert.hdf5 --epoch 12 --min-q-weight 50经验结论CRR 本身在 Atari 任务上表现不佳但叠加 CQL 损失/正则项后有明显改善Pong 从 -21 提升到 17.7Breakout 从 23.3 提升到 76.9。这一CRR w/ CQL模式通过 atari_crr.py 中的--min-q-weight参数实现——该参数控制 CQL 正则项的权重配合--policy_improvement_mode exp、--ratio_upper_bound 20.0、--beta 1.0等 CRR 自身超参数使用。使用 RL Unplugged 数据集除自采的 QRDQN 专家数据外Tianshou 还提供脚本 convert_rl_unplugged_atari.py用于将 DeepMindRL Unplugged的 Atari 数据集转换为 TianshouReplayBuffer可读的 HDF5 格式。数据转换命令以下命令会下载 Breakout 游戏第一次运行run 1的第一个分片shard 1转换为tianshou.data.ReplayBuffer并保存python3 convert_rl_unplugged_atari.py --task Breakout --run-id 1 --shard-id 1默认目录为~/.rl_unplugged/datasets/Breakout/run_1-00001-of-00100原始 TFRecord 缓存与~/.rl_unplugged/buffers/Breakout/run_1-00001-of-00100.hdf5转换结果可用--dataset-dir与--buffer-dir修改。脚本还支持--cache-dir指定下载缓存目录并会检查环境变量RLU_CACHE_DIR与RLU_DATASET_DIR。用转换后的数据训练python3 atari_bcq.py --task BreakoutNoFrameskip-v4 --load-buffer-name ~/.rl_unplugged/datasets/Breakout/run_1-00001-of-00100.hdf5 --buffer-from-rl-unplugged --epoch 12注意这里必须加--buffer-from-rl-unplugged开关脚本才会走 HDF5 加载路径atari_bcq.py 中对应if args.buffer_from_rl_unplugged: buffer load_buffer(...)的分支。转换脚本的实现细节从 convert_rl_unplugged_atari.py 的源码可以了解其工作原理数据规模RL Unplugged 每个 run 记录全部 5000 万条转移对应 2 亿环境步因 4 倍帧跳过每 100 个分片每个分片约含 50 万条转移每个游戏有 5 个独立 run转换流程process_shard从 Google Cloud Storage 下载分片 → 用tf.data.TFRecordDataset流式读取 TFRecordGZIP 压缩→ 逐条把tf.Example解析并解码 4 帧 PNG形状4 x 84 x 84→ 组装成Batch(obs, act, rew, done, obs_next)→ 按 D4RL 命名惯例写入 HDF5observations/actions/rewards/terminals/next_observations五个数据集gzip 压缩。其中done 1 - d_t对应 RL Unplugged 中的 discount 字段覆盖游戏支持 9 个调参游戏 36 个测试游戏共 45 个 Atari 游戏TUNING_SUITE与TESTING_SUITE列表。使用注意事项转换脚本依赖 Tensorflow仅用于解析 TFRecord 与解码 PNG脚本开头已通过tf.config.set_visible_devices([], GPU)禁用 GPU处理一个分片大约需要1 小时作者机器上的耗时实际情况因机器而异YMMV。离线 RL 示例的通用运行模式与调试技巧综合以上所有示例脚本可以总结出 Tianshou 离线 RL 示例的通用骨架便于读者复用环境与空间信息连续控制用SpaceInfo.from_env获取obs_shape/action_shape/max_action/min_action离散控制用make_atari_env构建 Atari 环境并读取observation_space.shape与action_space.n网络与算法组装按各算法的论文架构构建网络再封装为Policy与Algorithm如BCQPolicy/BCQ、SACPolicy/CQL、DiscreteBCQPolicy/DiscreteBCQ数据加载D4RL 走load_buffer_d4rl自采 Atari 数据走pickle/VectorReplayBuffer.load_hdf5RL Unplugged 数据走load_buffer--buffer-from-rl-unplugged训练统一调用algorithm.run_training(OfflineTrainerParams(buffer..., test_collector..., ...))日志与回放所有脚本支持--logger tensorboard|wandb与--watch回放模式训练后的最优策略由save_best_fn保存为policy.pth。调试方面以下几个参数最常需要调节连续控制中--epoch/--epoch_num_steps控制训练总量--batch_size控制每次梯度更新的样本量离散控制中--min-q-weightCQL/CRR 正则权重对最终性能影响显著数据量越小越需要降低该值参考 1e6/1e5/1e4 三档 buffer 下分别使用 50/20/10 的经验值--watch模式下需要配合--resume-path指定已保存的策略文件。小结Tianshou 的 examples/offline 目录提供了一个开箱即用的离线强化学习实验套件连续控制方向覆盖 D4RL 上的 IL、BCQ、CQL、TD3BC 四种算法离散控制方向覆盖 Atari 上的 IL、BCQ、CQL、CRR 四种算法并支持从自采专家数据到 RL Unplugged 标准数据集的完整数据管线。所有示例共享同一套OfflineTrainer训练机制通过ReplayBuffer统一数据接口配合文档中记录的可复现基准结果与超参数经验是快速开展离线 RL 实验与算法对比的实用起点。赞分享人工智能机器学习深度学习强化学习【免费下载链接】tianshouAn elegant PyTorch deep reinforcement learning library.项目地址https://gitcode.com/gh_mirrors/ti/tianshou点击查看免费下载相关推荐JRL 中的 CQL 离线强化学习实现配置参数、BC 预热与 D4RL 训练实战指南JRL 中的 CQL 离线强化学习实现配置参数、BC 预热与 D4RL 训练实战指南 本指南以 Google Research 的 JRLJax Reinf人工智能深度学习NLP计算机视觉强化学习Tianshou项目中的离线强化学习实践指南Tianshou项目中的离线强化学习实践指南 什么是离线强化学习 离线强化学习 Offline Reinforcement Learning 是一种特殊的强化学人工智能机器学习深度学习强化学习AReaL 数据集定制完全指南从 SFT 离线训练到 GRPO 在线强化学习AReaL 数据集定制完全指南从 SFT 离线训练到 GRPO 在线强化学习 AReaLThe RL Bridge for LLM based Agent人工智能大模型强化学习分布式训练AI Agent创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考