ARTICLE DETAIL

资讯详情

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

TensorTrade 并行 DQN 智能体(ParallelDQNAgent)源码级解析:多进程训练架构、参数详解与实战指南

TensorTrade 并行 DQN 智能体(ParallelDQNAgent)源码级解析:多进程训练架构、参数详解与实战指南 人工智能金融科技机器学习【免费下载链接】tensortradeAn open source reinforcement learning framework for training, evaluating, and deploying robust trading agents.项目地址https://gitcode.com/gh_mirrors/te/tensortrade点击查看免费下载本指南以docs/source/api/tensortrade.agents.parallel.rst所定义的tensortrade.agents.parallel包为骨架结合仓库源码深入讲解 TensorTrade 内置的并行 DQN 智能体它如何通过多个训练进程并行采集经验、一个优化进程异步更新网络以及n_envs、batch_size、eps_decay_steps、update_target_every等关键超参的作用与默认值。读完本文你将掌握ParallelDQNAgent的完整训练流程、消息队列通信机制、模型/优化器/训练器四个子模块的分工以及如何用它并行加速强化学习交易智能体的训练。1. 模块总览parallel 包在 TensorTrade 中的定位tensortrade.agents.parallel是 TensorTrade 内置智能体体系中的并行 DQN 实现。在tensortrade/agents/__init__.py中它与DQNAgent、A2CAgent一起被导出用户可直接通过from tensortrade.agents import ParallelDQNAgent引入。该包由五个模块组成对应docs/source/api/tensortrade.agents.parallel.rst中列出的五个子模块tensortrade.agents.parallel.parallel_dqn_agent、parallel_dqn_model、parallel_dqn_optimizer、parallel_dqn_trainer、parallel_queue顶层__init__.py将五个类全部暴露from .parallel_dqn_agent import ParallelDQNAgent from .parallel_dqn_model import ParallelDQNModel from .parallel_dqn_optimizer import ParallelDQNOptimizer from .parallel_dqn_trainer import ParallelDQNTrainer from .parallel_queue import ParallelQueue从架构上看四个核心类 一个基础设施类的职责划分如下类文件职责ParallelDQNAgentparallel_dqn_agent.py用户入口编排训练过程、启动各子进程、汇总结果ParallelDQNModelparallel_dqn_model.py策略网络与目标网络、动作选择、模型保存/恢复ParallelDQNTrainerparallel_dqn_trainer.py每个环境一个进程采样探索/利用、写入经验队列ParallelDQNOptimizerparallel_dqn_optimizer.py单独进程读取经验、采样 batch、梯度更新、回传新模型ParallelQueueparallel_queue.py可移植的多进程队列解决 Unix 上qsize()不可用问题1.1 设计要点数据生产者与消费者解耦ParallelDQNAgent继承自tensortrade.agents.agent.Agentagent.py必须实现restore、save、get_action、train四个抽象方法。其核心思想是把采样experience collection与学习gradient descent解耦为两类进程中间用共享内存级别的多进程队列传递数据从而让采样速度不再受反向传播瓶颈限制实现多环境并行采集。2. 训练架构Trainer × N 与 Optimizer × 1ParallelDQNAgent.train()parallel_dqn_agent.py#L106-L172的完整流程如下从**kwargs读取全部训练超参带默认值见第 3 节创建三条ParallelQueuememory_queue经验样本、model_update_queue模型更新、done_queue完成信号启动n_envs个ParallelDQNTrainer进程每个进程内部create_env()创建独立交易环境启动 1 个ParallelDQNOptimizer守护进程optimizer_process.daemon True主进程每 5 秒轮询一次done_queue直到所有 trainer 完成累加各 trainer 回报得到mean_reward total_reward / n_envs并返回依次关闭并join_thread三条队列terminatejoin所有 trainer 进程。其进程拓扑如下图所示文字版主进程 ParallelDQNAgent.train() ├── Trainer 进程 #1env #1──采样──▶ memory_queue ├── Trainer 进程 #2env #2──采样──▶ memory_queue ├── ...共 n_envs 个默认 CPU 核数 └── Optimizer 进程daemon◀──读取── memory_queue │ 从 model_update_queue 回传新模型 ├──▶ Trainer 们读取 model_update_queue 并同步网络 └──▶ done_queue 汇总结束信号2.1 Trainer 进程的循环逻辑每个ParallelDQNTrainerparallel_dqn_trainer.py#L54-L97在run()中执行探索/利用 采样循环每轮 episode 开始时若model_update_queue非空则取出最新的模型并调用self.agent.model.update_networks(model)同步策略网络与目标网络权重parallel_dqn_model.py#L80-L82使用指数衰减的 ε-greedy 策略计算当前动作阈值threshold eps_end (eps_start - eps_end) * np.exp(-steps_done / eps_decay_steps)执行env.step(action)得到(next_state, reward, done, _)将五元组(state, action, reward, next_state, done)推入memory_queue每update_target_every步调用一次update_target_network()把策略网络权重复制给目标网络parallel_dqn_model.py#L84-L86满足n_steps默认np.iinfo(np.int32).max即不限步数或n_episodes上限后退出将mean_reward total_reward / steps_done放入done_queue。注意Trainer 进程在构造时还会执行self.env.agent_id self.agent.id为每个训练环境绑定 agent 的唯一标识id由Identifiable基类生成见 core/base.py#L20-L46。2.2 Optimizer 进程的梯度更新ParallelDQNOptimizerparallel_dqn_optimizer.py#L49-L97在run()中循环执行只要done_queue中完成信号数量小于n_envs就持续工作把memory_queue中的样本全部取出并memory.push(*sample)写入本地ReplayMemory容量memory_capacity转移类型为DQNTransition见 replay_memory.py若本地经验不足一个 batchcontinue等待更多采样否则memory.sample(batch_size)随机采样用DQNTransition(*zip(*transitions))重组 batch在tf.GradientTape中计算 TD 目标当前状态动作值reduce_sum(policy_network(state_batch) * one_hot(action_batch), axis1)下一状态值对done样本置零否则取max(target_network(next_state_batch))期望值reward_batch discount_factor * next_state_values损失Huber 损失tf.keras.losses.Huber()优化器为 Nadamtf.keras.optimizers.Nadam计算梯度并apply_gradients随后把更新后的模型对象放入model_update_queue广播给各 Trainer。值得一提的是ParallelDQNOptimizer.__init__中默认learning_rate0.001注释中保留了 0.0001 的旧值而ParallelDQNAgent.train()传入的默认learning_rate0.0001两者存在默认值差异从源码结构看是开发演进过程中留下的不一致用户应显式传入learning_rate以避免歧义。3. 关键参数详解与默认值对照train()的全部超参均通过**kwargs传入parallel_dqn_agent.py#L113-L121参数默认值作用传递目标n_envsmp.cpu_count()并行训练环境Trainer 进程数量决定 Trainer 进程数batch_size128每次梯度更新的采样批次大小经验不足 batch 时跳过学习Optimizerdiscount_factor0.9999未来奖励折扣 γOptimizerTD 目标计算learning_rate0.0001Nadam 优化器学习率Optimizereps_start0.9ε-greedy 初始探索率Trainereps_end0.05ε-greedy 最低探索率Trainereps_decay_steps2000探索率指数衰减的时间尺度Trainer阈值公式update_target_every1000每多少步将策略网络权重同步到目标网络Trainermemory_capacity10000经验回放缓冲区容量OptimizerReplayMemorytrain()本身的命名参数包括n_steps、n_episodes、save_every、save_path、callback其中n_steps/n_episodes会原样传给 TrainerTrainer 内部用np.iinfo(np.int32).max作为缺省上限。save_every、save_path、callback在当前ParallelDQNAgent.train()实现中未被使用从源码结构看这是与DQNAgent.train()相比尚未完整落地的能力。3.1 ε 衰减曲线与update_target_every的协同Trainer 的探索阈值按eps_end (eps_start - eps_end) * exp(-steps_done / eps_decay_steps)单调衰减eps_decay_steps越大探索期越长。与此同时Trainer 每update_target_every步就复制一次策略权重到目标网络这与经典 DQN每 N 步软/硬更新目标网络一致但此处为硬复制。由于每个 Trainer 只更新自己的目标网络副本而策略网络权重统一由 Optimizer 通过model_update_queue广播因此所有 Trainer 最终会收敛到同一份策略。4. ParallelDQNModel网络结构与动作决策ParallelDQNModelparallel_dqn_model.py#L27-L86在构造时通过create_env()临时创建环境来探测action_space.n与observation_space.shape随后默认策略网络为tf.keras.Sequential_build_policy_networkInputLayer → Conv1D(64, k6, tanh) → MaxPooling1D(2) → Conv1D(32, k3, tanh) → MaxPooling1D(2) → Flatten → Dense(n_actions, sigmoid) → Dense(n_actions, softmax)target_network tf.keras.models.clone_model(policy_network)且trainable Falseget_action(state, threshold0)随机数小于 threshold 时均匀随机选动作探索否则取argmax(policy_network(expand_dims(state, 0)))利用save(path, agent_id, episode)按policy_network__{agent_id}__{episode:03d}.hdf5带 episode或policy_network__{agent_id}.hdf5命名保存restore(path)用load_model恢复并用新权重重建目标网络。这些能力经ParallelDQNAgent的save/restore/get_action/update_networks/update_target_network方法暴露给外部与Agent抽象基类的接口一一对应。5. ParallelQueue可移植的多进程队列ParallelQueueparallel_queue.py#L63-L95)继承multiprocessing.queues.Queue解决了一个具体痛点在 macOS 等平台上sem_getvalue()未实现会导致Queue.qsize()抛出NotImplementedError。实现方式是组合一个SharedCounterparallel_queue.py#L23-L60用mp.Value(i, 0)mp.Lock保证计数原子性注释中明确解释了n 1是读后写需要 Lock 保证原子性这一多进程经典问题put()时计数 1、get()时计数 -1从而提供可靠的qsize()与empty()。这正是train()主循环里while done_queue.qsize() n_envs轮询能够工作的前提。6. 使用示例与注意事项6.1 最小训练代码from tensortrade.agents import ParallelDQNAgent def create_env(): # 返回一个 tensortrade TradingEnvironment 实例 ... agent ParallelDQNAgent(create_envcreate_env) mean_reward agent.train( n_steps5000, # 每个 Trainer 最多执行步数 n_episodes20, # 每个 Trainer 最多 episode 数 n_envs4, # 并行环境数默认 CPU 核数 batch_size128, discount_factor0.9999, learning_rate0.0001, eps_start0.9, eps_end0.05, eps_decay_steps2000, update_target_every1000, memory_capacity10000, ) agent.save(agent/) # 生成 policy_network__id.hdf5 agent.restore(agent/policy_network__id.hdf5)ParallelDQNAgent的构造签名仅需create_env可调用对象每次调用返回一个全新的TradingEnvironment模型可选项modelParallelDQNModel(create_envcreate_env)。6.2 重要注意事项弃用状态从源码看ParallelDQNAgent、ParallelDQNModel、ParallelDQNTrainer、ParallelDQNOptimizer、ParallelQueue五个类均标注deprecated(version1.0.4, reasonBuiltin agents are being deprecated in favor of external implementations (ie: Ray))。项目正将内置智能体迁移到外部实现如 Ray RLlib本文内容基于当前仓库 1.0.5-dev 版本见 version.py的实际代码生产环境建议关注迁移方案。无对应单元测试当前仓库tests/目录中未发现针对 parallel 包的测试用例建议在自定义环境上先做小规模冒烟训练如n_envs2、少量步数验证队列与进程通信后再放大规模。进程安全create_env必须在每个 Trainer 子进程内被调用以创建独立环境因此传入的函数必须是可 pickle 的顶层函数或可构造对象避免在进程间共享不可序列化的环境状态。样本不再使用与DQNAgent的在线学习不同ParallelDQNAgent不执行env.render()/env.save()也不按 episode 自动保存 checkpointsave_every、save_path参数在当前实现中未生效需要用户手动调用agent.save(path)。7. 关联文档与 API 参考本文对应的 API 文档入口为 tensortrade.agents.parallel.rst其子模块文档页分别为tensortrade.agents.parallel.parallel_dqn_agent.rsttensortrade.agents.parallel.parallel_dqn_model.rsttensortrade.agents.parallel.parallel_dqn_optimizer.rsttensortrade.agents.parallel.parallel_dqn_trainer.rsttensortrade.agents.parallel.parallel_queue.rst这些页面均为 Sphinxautomodule自动生成的 API 参考内容即上述五个模块的成员、继承关系与文档字符串本文通过直接阅读 tensortrade/agents/parallel/ 下的源码对每个符号的行为进行了逐一印证与补充说明。若需深入了解单进程 DQN 基线可对照阅读 dqn_agent.py 与 agent.py以便在并行与串行训练间做公平对比。赞分享人工智能金融科技机器学习【免费下载链接】tensortradeAn open source reinforcement learning framework for training, evaluating, and deploying robust trading agents.项目地址https://gitcode.com/gh_mirrors/te/tensortrade点击查看免费下载相关推荐TensorTrade ParallelDQNAgent 深度解析基于多进程架构的并行 DQN 强化学习交易智能体TensorTrade ParallelDQNAgent 深度解析基于多进程架构的并行 DQN 强化学习交易智能体 导读 ParallelDQNAgent 是人工智能金融科技机器学习TensorTrade ParallelQueue 源码剖析为并行 DQN 训练定制可靠的多进程队列TensorTrade ParallelQueue 源码剖析为并行 DQN 训练定制可靠的多进程队列 导读 本文围绕 TensorTrade 仓库中的 API人工智能金融科技机器学习TensorTrade Parallel DQN 训练器ParallelDQNTrainer源码级解析多进程强化学习训练工作进程的完整实现TensorTrade Parallel DQN 训练器ParallelDQNTrainer源码级解析多进程强化学习训练工作进程的完整实现 导读 本文以人工智能金融科技机器学习上一篇DrissionPage在青龙面板中安装失败的解决方案下一篇APIPark项目中自定义AI模型添加问题的分析与解决创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表