ARTICLE DETAIL

资讯详情

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

TensorTrade Agent 抽象基类解析:从接口设计到 DQN / A2C 实战训练

TensorTrade Agent 抽象基类解析:从接口设计到 DQN / A2C 实战训练 人工智能金融科技机器学习【免费下载链接】tensortradeAn open source reinforcement learning framework for training, evaluating, and deploying robust trading agents.项目地址https://gitcode.com/gh_mirrors/te/tensortrade点击查看免费下载本文围绕 TensorTrade 开源强化学习框架中的tensortrade.agents.agent模块展开深入解析框架内智能体Agent的抽象接口设计、四个核心抽象方法restore、save、get_action、train的语义与调用约定并结合仓库内的DQNAgent、A2CAgent、ParallelDQNAgent实现与 Notebook 示例给出可直接运行的训练、保存与恢复实战方案。读完本文你将掌握如何基于 TensorTrade 自定义一个符合框架约定的强化学习智能体并理解内置 Agent 的弃用原因与迁移方向。一、Agent 在 TensorTrade 框架中的位置TensorTrade 是一个面向训练、评估与部署稳健交易智能体的开源强化学习框架。在框架的分层架构中tensortrade.agents.agent模块承担着学习智能体这一层环境Environment负责产生观测与奖励而 Agent 则负责把观测映射为动作并在训练循环中更新自身的策略模型。本项目中的 Agent 定义于 tensortrade/agents/agent.py并通过 tensortrade/agents/init.py 统一导出Agent、ReplayMemory、DQNAgent、DQNTransition、A2CAgent、A2CTransition以及ParallelDQNAgent。从框架的 API 文档组织看docs/source/api/tensortrade.agents.agent.rst正是通过 Sphinx 的automodule指令自动生成该模块的 API 参考页其核心内容即Agent抽象基类。二、Agent 抽象基类四个抽象方法构成的最小接口Agent继承自Identifiable见 tensortrade/core/base.py因此每个 Agent 实例都拥有一个通过uuid.uuid4()生成的唯一id属性用于在日志、模型文件名中标识不同的训练实例。deprecated(version1.0.4, reasonBuiltin agents are being deprecated in favor of external implementations (ie: Ray)) class Agent(Identifiable, metaclassABCMeta): abstractmethod def restore(self, path: str, **kwargs): ... abstractmethod def save(self, path: str, **kwargs): ... abstractmethod def get_action(self, state: np.ndarray, **kwargs) - int: ... abstractmethod def train(self, n_steps: int None, n_episodes: int 10000, save_every: int None, save_path: str None, callback: callable None, **kwargs) - float: ...四个抽象方法的约定如下抽象方法签名要点语义约定restorepath: str, **kwargs从path指定的文件恢复 Agent 的模型参数用于继续训练或部署推理savepath: str, **kwargs将 Agent 的模型保存到path指定目录通常配合episode参数生成带时间戳的检查点文件名get_actionstate: np.ndarray, **kwargs - int针对环境的某个具体状态返回一个动作索引整数是推理阶段的核心入口trainn_steps、n_episodes、save_every、save_path、callback、**kwargs - float在环境中训练 Agent最终返回训练期间的平均奖励mean reward值得注意的细节train方法的基类默认值为n_episodes10000、n_stepsNone但各具体实现会给出不同的默认值如 DQN 默认n_steps1000, n_episodes10因此阅读具体实现时需以子类签名为准。callback参数在基类签名中已预留但当前内置实现DQN / A2C / ParallelDQN的train方法中尚未实际调用它这是从源码结构推断出的预留接口。get_action返回的是int类型的动作索引与环境action_space中的离散动作一一对应。三、内置实现的源码级解读3.1 DQNAgent深度 Q 网络的完整实现DQNAgenttensortrade/agents/dqn_agent.py是 Agent 接口最典型的落地实现其训练过程遵循经典的 DQN 算法初始化__init__从环境读取n_actions int(env.action_space.n)与observation_shape env.observation_space.shape若未传入自定义policy_network则调用_build_policy_network()构建默认网络通过tf.keras.models.clone_model克隆出target_network目标网络并冻结为trainable False将自身id写入env.agent_id把智能体与环境绑定。默认策略网络结构输入层接收observation_shape的观测随后是 3 路并行的Conv1D因果卷积16/32/64 个滤波器kernel_size4strides2PReLU激活 he_uniform初始化经BatchNormalization后拼接再经过两层 Dropout(0.9)、第二轮三路卷积、AveragePooling1D、四层GRU(64)最后以Dense(n_actions, sigmoid) → Dense(n_actions, softmax)输出动作概率。该结构同时混合了卷积时序特征提取与 GRU 序列建模适合价格序列类观测。get_action与 ε-greedy 探索def get_action(self, state: np.ndarray, **kwargs) - int: threshold: float kwargs.get(threshold, 0) rand random.random() if rand threshold: return np.random.choice(self.n_actions) else: return np.argmax(self.policy_network(np.expand_dims(state, 0)))当随机数小于threshold时执行随机探索否则取策略网络输出最大值对应的动作。train主循环与超参数train通过kwargs读取一批可调超参数默认值与作用如下超参数DQN 默认值作用batch_size256每次梯度下降采样的样本数内存中样本不足时跳过更新memory_capacityn_steps * 10经验回放池容量discount_factor0.95未来奖励折扣因子 γlearning_rate0.01Nadam 优化器学习率eps_start/eps_end0.9 / 0.05ε-greedy 探索率上下限eps_decay_stepsn_stepsε 指数衰减的时间常数update_target_every1000每多少步同步一次目标网络render_intervaln_steps // 10渲染间隔步数训练循环要点每 episode 执行self.env.reset()在done之前循环env.step(action)ε 按eps_end (eps_start - eps_end) * exp(-total_steps_done / eps_decay_steps)指数衰减兼容 5 元组next_state, reward, terminated, truncated, _Gymnasium 风格与 4 元组next_state, reward, done, _两种step返回值经验存入ReplayMemory当len(memory) batch_size时调用_apply_gradient_descent梯度更新使用Nadam优化器 Huber损失目标值由目标网络计算reward γ * max Q(next_state)done状态的目标值置零每update_target_every步重新克隆策略网络为目标网络满足save_every间隔或最后一个 episode 时调用save保存检查点每 episode 结束调用self.env.save()最终返回total_reward / steps_done作为平均奖励。模型持久化约定save生成的文件名为policy_network__{agent_id前7位}__{YYYYmmdd_HHMMSS}.hdf5保存到path filenamerestore则用tf.keras.models.load_model(path)加载并重建目标网络。3.2 A2CAgentActor-Critic 双网络结构A2CAgenttensortrade/agents/a2c_agent.py实现 Advantage Actor-CriticA2C算法其网络结构与 DQN 截然不同shared_network共享特征提取层Conv1D(64, 6) → MaxPooling1D(2) → Conv1D(32, 3) → MaxPooling1D(2) → Flattenactor_network共享网络 Dense(50, relu) → Dense(n_actions, relu)输出动作 logitscritic_network共享网络 Dense(50) → Dense(25) → Dense(1, relu)输出状态价值。get_action使用tf.random.categorical从 actor 输出的 logits 中采样动作配合threshold实现随机探索。_apply_gradient_descent的两阶段更新从记忆尾部取出batch_size条经验按时间倒序计算折扣回报returnsreward γ * return * (1 - done)Critic 用Huber损失拟合回报用Adam优化器更新Actor 用SparseCategoricalCrossentropy(from_logitsTrue)并以优势值returns - values作为样本权重计算策略损失再减去熵正则项entropy_c * entropy鼓励探索。A2C 关键超参数默认值batch_size128、discount_factor0.9999、learning_rate0.0001、eps_decay_steps200、entropy_c0.0001、memory_capacity1000。持久化差异A2C 分别保存actor_network__...hdf5与critic_network__...hdf5两个文件restore必须同时传入actor_filename与critic_filename两个kwargs否则抛出ValueError。3.3 ParallelDQNAgent多进程并行训练ParallelDQNAgenttensortrade/agents/parallel/parallel_dqn_agent.py接受一个create_env工厂函数而非环境实例通过multiprocessing并行启动n_envs默认mp.cpu_count()个ParallelDQNTrainer进程各自构建环境、采集经验并写入memory_queue一个守护进程ParallelDQNOptimizer从队列消费经验、更新模型并通过model_update_queue把新权重回传done_queue汇总各环境的累计奖励最终返回total_reward / n_envs的平均奖励。该实现将采样与优化解耦为独立的进程队列是从源码结构推断出的并行化设计意图。其训练参数默认值batch_size128、discount_factor0.9999、learning_rate0.0001、eps_decay_steps2000、update_target_every1000、memory_capacity10000。3.4 ReplayMemory经验回放的基础设施ReplayMemorytensortrade/agents/replay_memory.py是 DQN / A2C 共同依赖的环形缓冲push(*args)按环形队列覆盖写入容量为capacitysample(batch_size)随机采样DQN 使用head(batch_size)/tail(batch_size)分别取头部/尾部连续切片A2C 使用tail以便按时间倒序计算回报过渡元组类型通过transition_type参数注入如DQNTransition(state, action, reward, next_state, done)与A2CTransition(state, action, reward, done, value)。四、实战基于 Notebook 的完整训练流程仓库中的 examples/train_and_evaluate.ipynb 给出了DQNAgent端到端训练的完整代码。以下是核心流程源码可复现第 1 步构建数据、交易所与投资组合import tensortrade.env.default as default from tensortrade.agents import DQNAgent from tensortrade.feed.core import DataFeed, Stream from tensortrade.env.default.actions import BSH from tensortrade.env.default.rewards import RiskAdjustedReturns, SimpleProfit from tensortrade.oms.exchanges import Exchange from tensortrade.oms.services.execution.simulated import execute_order from tensortrade.oms.instruments import USD, BTC from tensortrade.oms.wallets import Wallet, Portfolio price Stream.source(list(X_train[close]), dtypefloat).rename(USD-BTC) bitstamp Exchange(bitstamp, serviceexecute_order)(price) cash Wallet(bitstamp, 50000 * USD) asset Wallet(bitstamp, 0 * BTC)第 2 步组装 Feed、动作方案与奖励方案创建环境feed DataFeed([price, price.rolling(10).mean().rename(fast), ...]) reward_scheme RiskAdjustedReturns() # 或 SimpleProfit() action_scheme BSH(cashcash, assetasset).attach(reward_scheme) env default.create(feedfeed, portfolioportfolio, action_schemeaction_scheme, reward_schemereward_scheme, window_sizewindow_size, max_allowed_loss0.6)第 3 步实例化 DQNAgent 并训练agent DQNAgent(env) agent.train(batch_sizebatch_size, n_stepsn_steps, n_episodesn_episodes, memory_capacitymemory_capacity, save_pathsave_path)训练完成后可通过agent.save(path)持久化模型在后续会话中用agent.restore(path)恢复并调用agent.get_action(state)进行推理。说明Notebook 中batch_size等参数由辅助函数get_optimal_batch_size(window_sizewindow_size, n_stepsn_steps, batch_factor4)计算得出实际运行时可根据 3.1 节的默认值表自行指定。五、兼容性与迁移方向内置 Agent 已弃用从源码可见Agent、ReplayMemory、DQNAgent、A2CAgent、ParallelDQNAgent均带有deprecated(version1.0.4, reasonBuiltin agents are being deprecated in favor of external implementations (ie: Ray))装饰器。这意味着自 1.0.4 版本起内置 Agent 被标记为弃用官方推荐迁移到外部强化学习实现典型代表是 Ray 的 RLLib在迁移指南 MIGRATION_GUIDE.md 中Agent framework 被明确列为向后兼容的部分即既有基于Agent接口编写的自定义组件仍可继续使用弃用但未移除tensortrade/agents/__init__.py仍完整导出上述类历史代码可继续运行但新项目应优先采用外部 RL 框架。docs/source/agents/overview.md展示了替代方向通过ray.tune.run(PPO, ...)训练策略、用ray.rllib.agents.ppo.PPOTrainer恢复检查点以及 Tensorforce、Stable Baselines 等库的接入方式。这也印证了TensorTrade 框架本身与多种强化学习库互操作的设计目标——Agent 抽象层正是这种互操作性的边界自定义 Agent 只需实现restore/save/get_action/train四个方法即可无缝接入 TensorTrade 的环境、OMS 与数据流体系。六、自定义 Agent 的推荐实践基于Agent抽象类编写一个自定义智能体应遵循以下步骤继承Agent会自动获得Identifiable.id能力并在构造函数中把self.id写入env.agent_id实现get_action(state, threshold...)返回int动作索引保留 ε 阈值参数以兼容现有训练循环的调用方式实现train(n_steps, n_episodes, save_every, save_path, callback, **kwargs)内部通过kwargs.get(key, default)读取全部可调超参数循环env.reset()/env.step()并在 checkpoint 时机调用self.save最终返回平均奖励float实现save(path)与restore(path)建议沿用{network_name}__{agent_id[:7]}__{timestamp}.hdf5的命名约定便于多 Agent 并行训练时区分检查点经验缓存可复用ReplayMemory通过transition_type注入自定义namedtuple过渡类型。七、总结tensortrade.agents.agent模块以 4 个抽象方法定义了 TensorTrade 学习智能体的最小契约restore负责恢复、save负责持久化、get_action负责策略推理、train负责训练循环并返回平均奖励。围绕这一契约仓库提供了 DQN含目标网络、经验回放、ε-greedy 衰减、A2CActor-Critic 双网络 熵正则与多进程并行 DQN 三套完整实现并有 train_and_evaluate.ipynb 提供开箱即用的实战示例。由于内置 Agent 自 1.0.4 起被弃用新项目建议基于该抽象接口自行实现或直接迁移到 Ray RLLib 等外部强化学习框架而理解这套接口约定正是自定义智能体、复用 TensorTrade 环境与 OMS 能力的前提。赞分享人工智能金融科技机器学习【免费下载链接】tensortradeAn open source reinforcement learning framework for training, evaluating, and deploying robust trading agents.项目地址https://gitcode.com/gh_mirrors/te/tensortrade点击查看免费下载相关推荐sngrep源码解析从packet捕获到UI渲染的完整技术流程sngrep源码解析从packet捕获到UI渲染的完整技术流程 sngrep是一款基于Ncurses的SIP消息流查看工具能够实时捕获、解析和可视化SIP协网络与通信运维EOS抽象基类ABC设计与接口规范化实践EOS抽象基类ABC设计与接口规范化实践 概述 在能源优化系统Energy Optimization SystemEOS的开发过程中抽象基类Abst后端智能家居SMAT/ArkAnalyzer-HapRay性能分析器基类抽象接口设计与实现SMAT/ArkAnalyzer HapRay性能分析器基类抽象接口设计与实现 引言性能分析框架的核心基石 在OpenHarmony应用性能优化领域一个设开发工具性能测试移动开发OpenHarmony上一篇BilibiliVideoDownload故障排查指南从登录失败到下载中断的全面解决方案下一篇WebdriverIO 测试安全实践指南敏感数据遮蔽、日志脱敏与密钥防护创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表