ARTICLE DETAIL

资讯详情

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

TensorTrade 与 Ray RLlib 深度实践:分布式强化学习交易智能体的配置、训练与评估

TensorTrade 与 Ray RLlib 深度实践:分布式强化学习交易智能体的配置、训练与评估 人工智能金融科技机器学习【免费下载链接】tensortradeAn open source reinforcement learning framework for training, evaluating, and deploying robust trading agents.项目地址https://gitcode.com/gh_mirrors/te/tensortrade点击查看免费下载导读Ray RLlib 是 TensorTrade 训练流程中承担分布式强化学习任务的引擎。本文以docs/tutorials/04-training/02-ray-rllib.md为核心系统讲解如何在 TensorTrade 中初始化 Ray、注册交易环境、配置 PPO 算法、编写自定义回调以跟踪账户盈亏以及完成分布式训练、模型保存/加载与最终评估。读完本文你将掌握一套可直接落地的 RLlib TensorTrade 交易训练方案并理解其底层调用链与仓库内的真实工程实现。一、为什么交易训练选择 Ray RLlibRay RLlib 是一套可扩展的分布式强化学习库与 TensorTrade 的对接点集中在四方面算法覆盖广内置 PPO、DQN、A2C、SAC 等主流算法可直接通过对应的*Config类配置无需自行实现策略网络与训练循环原生分布式支持多 CPU worker 并行采样、GPU 训练以及多节点 Ray 集群横向扩展回调机制通过DefaultCallbacks子类在 episode 边界插入自定义逻辑用于采集交易盈亏、持仓比例等 TensorTrade 特有的业务指标Gym 兼容TensorTrade 的default.create构建的环境符合 Gym 观测/动作接口RLlib 可直接驱动。从仓库源码看tests/tensortrade/integration/rllib/test_ray_training.py已为 PPO 的单次迭代、LSTM 模型和 AttentionNet 模型分别建立了最小化验证用例examples/training/目录下的train_ray_long.py、train_best.py、run_ray_simulation.py等脚本则提供了完整的实战参考。二、环境准备与依赖安装运行 RLlib 训练需要以下核心依赖见 examples/requirements.txtpip install -r examples/requirements.txt其中与训练直接相关的关键项为ray[default,tune,rllib,serve]2.10.0,3.0分布式训练与调参基础设施torch2.0.0默认的神经网络框架配合.framework(torch)使用optuna3.0.0用于后续超参数自动优化见 04-training/03-optuna.md。TensorTrade 本体位于tensortrade/目录训练脚本直接导入其feed、oms、env三个核心子包。三、Ray 初始化与自定义环境注册3.1 初始化 Rayimport ray from ray.tune.registry import register_env from ray.rllib.algorithms.ppo import PPOConfig # 本地模式初始化 ray.init( num_cpus6, # 指定可用 CPU 核心数 ignore_reinit_errorTrue, # 已初始化时不报错 log_to_driverFalse # 降低日志冗长度 ) # 注册自定义环境工厂 register_env(TradingEnv, create_env)关键点说明num_cpus并非固定上限而是 Ray 调度器可用的 CPU 资源声明实际并行度还取决于num_env_runnersignore_reinit_errorTrue在 Jupyter 等多次执行场景中非常实用避免重复ray.init抛错仓库脚本习惯在结尾调用ray.shutdown()见 examples/training/train_ray_long.py单进程测试中则使用 pytest fixture 统一管理生命周期见 tests/tensortrade/integration/rllib/conftest.py。3.2 环境工厂函数RLlib 要求环境以工厂函数形式注册原因是每个 worker 进程都会调用该函数构建一份独立的环境副本。工厂函数必须做到每次调用都从零创建状态否则分布式并行采样会出现数据串扰。def create_env(config: Dict[str, Any]): Factory function for TradingEnv. data pd.read_csv(config[csv_filename]) price Stream.source(list(data[close]), dtypefloat).rename(USD-BTC) exchange Exchange( exchange, serviceexecute_order, optionsExchangeOptions(commissionconfig.get(commission, 0.001)) )(price) cash Wallet(exchange, config.get(initial_cash, 10000) * USD) asset Wallet(exchange, 0 * BTC) portfolio Portfolio(USD, [cash, asset]) features [Stream.source(list(data[c]), dtypefloat).rename(c) for c in config.get(feature_cols, [])] feed DataFeed(features) feed.compile() reward_scheme PBR(priceprice) action_scheme BSH(cashcash, assetasset).attach(reward_scheme) env default.create( feedfeed, portfolioportfolio, action_schemeaction_scheme, reward_schemereward_scheme, window_sizeconfig.get(window_size, 10), max_allowed_lossconfig.get(max_allowed_loss, 0.5) ) env.portfolio portfolio # 供回调访问 return env这里展示的组件全部来自仓库核心实现Exchange、ExchangeOptions来自 tensortrade/oms/exchanges/exchange.pycommission通过ExchangeOptions注入用于模拟交易成本Wallet、Portfolio来自 tensortrade/oms/wallets/env.portfolio portfolio这一行是关键技巧——它将 Portfolio 挂到环境上便于回调中读取net_worthPBRPosition-Based Returns基于持仓的收益与BSHBuy/Sell/Hold买入/卖出/持有分别来自 tensortrade/env/default/rewards.py 和 tensortrade/env/default/actions.pydefault.create最终组装出符合 Gym 接口的环境见 tensortrade/env/default/init.py。注意run_ray_simulation.py中的create_env还在读取 CSV 后做了.bfill().ffill()空值填充并在NameSpace上下文内构建特征流避免不同 worker 的流命名冲突这是多 worker 场景下的一个实用细节。四、PPO 配置从默认值到交易专用参数4.1 完整配置示例config ( PPOConfig() # 环境 .environment( envTradingEnv, env_config{ csv_filename: /path/to/data.csv, feature_cols: [ret_1h, rsi, trend], window_size: 17, max_allowed_loss: 0.32, commission: 0.003, initial_cash: 10000, } ) # 框架 .framework(torch) # 采样 worker .env_runners( num_env_runners4, # 并行环境数 rollout_fragment_length200, ) # 回调 .callbacks(MyCallbacks) # 训练超参数 .training( lr3.29e-05, gamma0.992, lambda_0.9, clip_param0.123, entropy_coeff0.015, vf_clip_param100.0, train_batch_size2000, sgd_minibatch_size256, num_sgd_iter7, model{ fcnet_hiddens: [128, 128], fcnet_activation: tanh, }, ) # 资源 .resources( num_gpus0, # 需要 GPU 训练时设为 1 ) )该配置与仓库中 examples/training/train_best.py 的最佳配置高度一致后者正是教程 04-training/01-first-training.md 所述的train_best.py实现其超参数来自 Optuna 100 轮试验。一个兼容性注意点仓库脚本如train_best.py在PPOConfig()后追加了.api_stack(enable_rl_module_and_learnerFalse, enable_env_runner_and_connector_v2False)以使用传统 API 栈保证与当前代码的稳定兼容实际使用时请按所用 Ray 版本核对。4.2 学习参数参数含义示例值选择理由lr学习率3.29e-05极低学习率保证训练稳定避免策略剧烈震荡gamma折扣因子0.992高折扣因子让智能体更看重长期收益lambda_GAE广义优势估计参数0.9在偏差与方差之间取平衡4.3 PPO 专属参数参数含义示例值选择理由clip_param策略更新幅度限制0.123中等裁剪幅度兼顾收敛速度与稳定性entropy_coeff熵奖励系数探索程度0.015低熵 更偏利用减少无效随机交易vf_clip_param价值函数裁剪100.0极大值 基本不裁剪价值函数4.4 批量训练参数参数含义示例值选择理由train_batch_size每次更新使用的样本数2000中等批量兼顾速度与稳定sgd_minibatch_size小批量大小256标准小批量尺寸num_sgd_iter每个批次上的 SGD 轮数7多轮内循环提升样本利用率版本说明较新的 RLlib 将sgd_minibatch_size/num_sgd_iter更名为minibatch_size/num_epochs仓库脚本即采用新命名见 examples/training/train_ray_long.py两套名称在配置层面含义对应请按安装版本选择。4.5 网络结构model字典控制策略与价值网络fcnet_hiddens: [128, 128]两层各 128 个神经元的全连接网络fcnet_activation: tanhtanh 激活函数输出有界适合价格类特征若换 LSTM{use_lstm: True, lstm_cell_size: 64}换 AttentionNet 则需配置use_attention及 transformer 相关参数二者均已在 tests/tensortrade/integration/rllib/test_ray_training.py 中有可运行的初始化用例。五、自定义回调在 episode 边界采集盈亏指标回调是把 TensorTrade 的Portfolio.net_worth变成 RLlib 训练指标的唯一桥梁。RLlib 每个 worker 里的base_env可能同时运行多个子环境sub-environments回调通过env_index定位当前 episode 对应的那个环境。from ray.rllib.algorithms.callbacks import DefaultCallbacks class TradingCallbacks(DefaultCallbacks): def on_episode_start(self, *, worker, base_env, policies, episode, env_indexNone, **kwargs): 每个 episode 开始时记录初始净值。 env base_env.get_sub_environments()[env_index] if hasattr(env, portfolio): episode.user_data[initial_worth] float(env.portfolio.net_worth) def on_episode_end(self, *, worker, base_env, policies, episode, env_indexNone, **kwargs): 每个 episode 结束时计算并上报盈亏指标。 env base_env.get_sub_environments()[env_index] if hasattr(env, portfolio): final_worth float(env.portfolio.net_worth) initial_worth episode.user_data.get(initial_worth, 10000) # 自定义指标 pnl final_worth - initial_worth pnl_pct (pnl / initial_worth) * 100 episode.custom_metrics[pnl] pnl episode.custom_metrics[pnl_pct] pnl_pct episode.custom_metrics[final_worth] final_worth使用方法同样是链式配置config ( PPOConfig() ... .callbacks(TradingCallbacks) )仓库中的train_best.py、train_ray_long.py均实现了同构回调其中train_ray_long.py的WalletTrackingCallbacks还同时写入episode.hist_data可保留净值随训练推进的历史序列。读取指标result algo.train() # 回调产生的指标 custom_metrics result.get(env_runners, {}).get(custom_metrics, {}) avg_pnl custom_metrics.get(pnl_mean, 0) avg_pnl_pct custom_metrics.get(pnl_pct_mean, 0) print(fAverage PL: ${avg_pnl:,.0f} ({avg_pnl_pct:.1f}%))RLlib 会对每个custom_metrics自动聚合出_mean、_min、_max等统计量训练脚本只需读取pnl_mean即可得到当前迭代的平均盈亏。六、训练循环手动验证与内置评估6.1 基础训练循环algo config.build() for i in range(100): result algo.train() reward result.get(env_runners, {}).get(episode_reward_mean, 0) custom result.get(env_runners, {}).get(custom_metrics, {}) pnl custom.get(pnl_mean, 0) print(fIter {i1}: Reward {reward:.1f}, PL ${pnl:,.0f})6.2 带验证集的训练循环金融时序存在非平稳性训练集上的 reward 上升不代表真实盈利能力因此仓库脚本普遍采用每 N 轮在独立验证集上评估仅在验证盈亏创新高时保存模型的策略algo config.build() best_val_pnl float(-inf) for i in range(100): result algo.train() if (i 1) % 10 0: val_pnl evaluate(algo, val_data) if val_pnl best_val_pnl: best_val_pnl val_pnl algo.save(/tmp/best_model) print(fIter {i1}: Val ${val_pnl:,.0f} *NEW BEST*) else: print(fIter {i1}: Val ${val_pnl:,.0f}) # 加载最佳模型用于测试 algo.restore(/tmp/best_model)这正是 examples/training/train_best.py 中main()的实际逻辑每 10 轮调用evaluate()在验证集上跑 10 个 episode最优时algo.save(/tmp/best_model)训练结束后algo.restore(/tmp/best_model)再用测试集做多档佣金水平回测。6.3 手动评估函数def evaluate(algo, data: pd.DataFrame, n_episodes: int 10) - float: 运行 n 个 episode返回平均盈亏。 csv_path /tmp/eval.csv data.to_csv(csv_path, indexFalse) env_config { csv_filename: csv_path, feature_cols: feature_cols, # ... 其余配置 } pnls [] for _ in range(n_episodes): env create_env(env_config) obs, _ env.reset() done truncated False while not done and not truncated: action algo.compute_single_action(obs) obs, _, done, truncated, _ env.step(action) pnl env.portfolio.net_worth - 10000 pnls.append(pnl) return np.mean(pnls)该函数在 examples/training/train_best.py 与 examples/training/train_optuna.py 中以evaluate(algo, data, feature_cols, config, n...)的形式出现核心不变写临时 CSV → 由create_env重建环境 → 用compute_single_action逐步推理 → 取portfolio.net_worth与初始现金的差值。6.4 使用 RLlib 内置评估config ( PPOConfig() ... .evaluation( evaluation_interval10, # 每 10 轮评估一次 evaluation_num_episodes5, evaluation_config{ env_config: val_env_config, } ) ) result algo.train() eval_reward result.get(evaluation, {}).get(episode_reward_mean, 0)run_ray_simulation.py即采用内置评估.evaluation(evaluation_interval2, evaluation_config{env_config: env_config_evaluation, explore: False})其中explore: False确保评估时关闭探索纯粹利用当前策略。内置评估适合快速验证手动评估则便于在验证集上附加业务指标如多档佣金回测。七、分布式训练多 CPU、GPU 与集群7.1 多 CPU 并行ray.init(num_cpus16) config ( PPOConfig() ... .env_runners(num_env_runners8) # 8 个并行 worker )每个 worker 持有独立的环境副本RLlib 自动在它们之间分配采样任务。num_env_runners与num_cpus的经验关系是worker 数应小于等于可用核心数并预留主进程与采样线程的资源。仓库脚本中train_best.py用ray.init(num_cpus6)num_env_runners4train_ray_long.py用num_cpus8num_env_runners4。7.2 GPU 训练config ( PPOConfig() ... .resources(num_gpus1) # 策略网络使用 1 张 GPU )num_gpus控制学习器Learner/Driver使用 GPU 的数量若显存充足可配合num_gpus_per_env_runner让采样 worker 也使用 GPU。7.3 集群训练# 连接 Ray 集群 ray.init(addressauto) config ( PPOConfig() ... .env_runners(num_env_runners32) # worker 分布在集群各节点 )addressauto会自动发现已启动的 Ray 集群典型启动方式为ray start --head与ray start --addresshead_ip:6379随后 worker 会按集群资源自动调度。需要在多机环境下确认各节点 Python 环境与 TensorTrade 包版本一致并确保数据 CSV 对每个 worker 节点可访问或通过ray.put分发。八、模型保存、加载与导出8.1 保存检查点# 保存检查点返回实际路径 checkpoint_path algo.save(/tmp/checkpoints) print(fSaved to: {checkpoint_path}) # 指定自定义名称保存 checkpoint_path algo.save(/tmp/my_model)8.2 加载模型# 从检查点恢复含迭代编号 algo.restore(/tmp/checkpoints/checkpoint_000050) # 或新建算法实例后恢复 from ray.rllib.algorithms.ppo import PPO algo PPO(configconfig) algo.restore(/tmp/my_model)注意restore要求当前config与保存时的环境注册、网络结构一致最稳妥的用法是先config.build()再restore。8.3 导出策略# 导出为 ONNX 用于部署 policy algo.get_policy() policy.export_model(/tmp/model_onnx)export_model将策略网络导出为 ONNX 格式便于脱离 Ray 环境做在线推理部署。九、常见问题排查9.1 内存不足.training( train_batch_size1000, # 调小批量 sgd_minibatch_size128, ) .env_runners(num_env_runners2) # 减少 worker9.2 训练缓慢.env_runners(num_env_runners8) # 增加并行 worker .resources(num_gpus1) # 或启用 GPU9.3 训练出现 NaN.training(lr1e-5) # 降低学习率PPO 默认已启用梯度裁剪NaN 多由学习率过高、特征含 NaN 或极端 reward 造成。教程 04-training/01-first-training.md 同时建议对特征列做.bfill().ffill()并用assert not data[feature_cols].isna().any().any()前置校验。9.4 环境无法正确重置def create_env(config): # 每次调用都全新创建 price Stream.source(...) # 新流 exchange Exchange(...) # 新交易所 # ...RLlib 会在每次 episode 结束时调用reset()若工厂函数复用了模块级可变状态会导致 episode 间数据污染。务必保证每次调用都重建 Stream、Exchange、Portfolio 等全部对象。9.5 episode 一启动即结束max_allowed_loss过小会在开局就触发止损终止。仓库默认值在 0.40.9 之间train_best.py用 0.32train_ray_long.py用 0.5run_ray_simulation.py用 0.9训练初期可适当放宽。十、替换 PPODQN、A2C 与 SACTensorTrade 默认环境的动作空间由BSH买/卖/持有产生为离散动作因此连续动作算法SAC需要自定义动作方案才能直接使用。10.1 DQN离散动作from ray.rllib.algorithms.dqn import DQNConfig config ( DQNConfig() .environment(envTradingEnv, env_configenv_config) .training( lr1e-4, gamma0.99, replay_buffer_config{capacity: 100000}, ) )10.2 A2C比 PPO 更简单的 on-policy 算法from ray.rllib.algorithms.a3c import A2CConfig config ( A2CConfig() .environment(envTradingEnv, env_configenv_config) .training( lr5e-4, gamma0.99, ) )10.3 SAC连续动作from ray.rllib.algorithms.sac import SACConfig # 需要连续动作空间需替换 BSH 动作方案 config ( SACConfig() .environment(envTradingEnv, env_configenv_config) .training( lr3e-4, gamma0.99, ) )仓库还提供了 RLlib 的 LSTM / AttentionNet 网络集成示例examples/use_lstm_rllib.ipynb、examples/use_attentionnet_rllib.ipynb以及在test_ray_training.py中验证过的模型配置可在需要时序建模时参考。十一、要点总结与进阶路径回顾本文核心结论分布式由 RLlib 托管只需设置num_env_runners采样、聚合、更新自动并行化环境工厂是成败关键必须每次调用都创建全新环境env.portfolio portfolio是回调读取净值的前提回调承载业务指标PL、盈亏百分比、交易次数等全部通过episode.custom_metrics上报PPO 是默认起点与 TensorTrade 离散动作BSH天然匹配默认超参数即来自 Optuna 调优务必使用独立验证集训练 reward 上升而验证盈亏回落就是过拟合信号。下一步可进入 04-training/03-optuna.md结合 examples/training/train_optuna.py 了解如何用 Optuna 自动搜索超参数完整的实验结果汇总见 docs/EXPERIMENTS.md。赞分享人工智能金融科技机器学习【免费下载链接】tensortradeAn open source reinforcement learning framework for training, evaluating, and deploying robust trading agents.项目地址https://gitcode.com/gh_mirrors/te/tensortrade点击查看免费下载相关推荐dragula 拖拽核心概念入门容器、drake 实例与选项体系完全解析dragula 拖拽核心概念入门容器、drake 实例与选项体系完全解析 dragula 是一款体积极小的浏览器拖拽库drag and drop它的口号人工智能金融科技机器学习TensorTrade 学习代理实战使用 Ray RLlib / Stable Baselines / Tensorforce 训练交易智能体TensorTrade 学习代理实战使用 Ray RLlib / Stable Baselines / Tensorforce 训练交易智能体 本文围绕 Te人工智能金融科技机器学习TensorTrade 入门指南用强化学习训练、评估与部署量化交易智能体TensorTrade 入门指南用强化学习训练、评估与部署量化交易智能体 本篇技术指南以 TensorTrade 开源仓库的 README 为主线系统讲解该人工智能金融科技机器学习上一篇Citra模拟器终极解决方案5步快速修复常见问题指南下一篇告别卡顿与模糊GLFW视频模式让显示器性能全开创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表