☰
TensorTrade 与 Ray RLlib 深度实践:分布式强化学习交易智能体的配置、训练与评估
2026/10/8 1:25:36 网站建设 项目流程
  • 人工智能
  • 金融科技
  • 机器学习

【免费下载链接】tensortrade

An 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 RLlib

Ray 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.txt):

pip install -r examples/requirements.txt

其中与训练直接相关的关键项为:

  • ray[default,tune,rllib,serve]>=2.10.0,<3.0:分布式训练与调参基础设施;
  • torch>=2.0.0:默认的神经网络框架(配合.framework("torch")使用);
  • optuna>=3.0.0:用于后续超参数自动优化(见 04-training/03-optuna.md)。

TensorTrade 本体位于tensortrade/目录,训练脚本直接导入其feed、oms、env三个核心子包。

三、Ray 初始化与自定义环境注册

3.1 初始化 Ray

import ray from ray.tune.registry import register_env from ray.rllib.algorithms.ppo import PPOConfig # 本地模式初始化 ray.init( num_cpus=6, # 指定可用 CPU 核心数 ignore_reinit_error=True, # 已初始化时不报错 log_to_driver=False # 降低日志冗长度 ) # 注册自定义环境工厂 register_env("TradingEnv", create_env)

关键点说明:

  • num_cpus并非固定上限,而是 Ray 调度器可用的 CPU 资源声明;实际并行度还取决于num_env_runners;
  • ignore_reinit_error=True在 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"]), dtype="float").rename("USD-BTC") exchange = Exchange( "exchange", service=execute_order, options=ExchangeOptions(commission=config.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]), dtype="float").rename(c) for c in config.get("feature_cols", [])] feed = DataFeed(features) feed.compile() reward_scheme = PBR(price=price) action_scheme = BSH(cash=cash, asset=asset).attach(reward_scheme) env = default.create( feed=feed, portfolio=portfolio, action_scheme=action_scheme, reward_scheme=reward_scheme, window_size=config.get("window_size", 10), max_allowed_loss=config.get("max_allowed_loss", 0.5) ) env.portfolio = portfolio # 供回调访问 return env

这里展示的组件全部来自仓库核心实现:

  • Exchange、ExchangeOptions来自 tensortrade/oms/exchanges/exchange.py,commission通过ExchangeOptions注入,用于模拟交易成本;
  • Wallet、Portfolio来自 tensortrade/oms/wallets/,env.portfolio = portfolio这一行是关键技巧——它将 Portfolio 挂到环境上,便于回调中读取net_worth;
  • PBR(Position-Based Returns,基于持仓的收益)与BSH(Buy/Sell/Hold,买入/卖出/持有)分别来自 tensortrade/env/default/rewards.py 和 tensortrade/env/default/actions.py;
  • default.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( env="TradingEnv", 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_runners=4, # 并行环境数 rollout_fragment_length=200, ) # 回调 .callbacks(MyCallbacks) # 训练超参数 .training( lr=3.29e-05, gamma=0.992, lambda_=0.9, clip_param=0.123, entropy_coeff=0.015, vf_clip_param=100.0, train_batch_size=2000, sgd_minibatch_size=256, num_sgd_iter=7, model={ "fcnet_hiddens": [128, 128], "fcnet_activation": "tanh", }, ) # 资源 .resources( num_gpus=0, # 需要 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_learner=False, enable_env_runner_and_connector_v2=False),以使用传统 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: "tanh":tanh 激活函数,输出有界,适合价格类特征;
  • 若换 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_index=None, **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_index=None, **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(f"Average P&L: ${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(f"Iter {i+1}: Reward {reward:.1f}, P&L ${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(f"Iter {i+1}: Val ${val_pnl:+,.0f} *NEW BEST*") else: print(f"Iter {i+1}: 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, index=False) 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_interval=10, # 每 10 轮评估一次 evaluation_num_episodes=5, 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_interval=2, evaluation_config={"env_config": env_config_evaluation, "explore": False}),其中"explore": False确保评估时关闭探索,纯粹利用当前策略。内置评估适合快速验证;手动评估则便于在验证集上附加业务指标(如多档佣金回测)。

七、分布式训练:多 CPU、GPU 与集群

7.1 多 CPU 并行

ray.init(num_cpus=16) config = ( PPOConfig() ... .env_runners(num_env_runners=8) # 8 个并行 worker )

每个 worker 持有独立的环境副本,RLlib 自动在它们之间分配采样任务。num_env_runners与num_cpus的经验关系是:worker 数应小于等于可用核心数,并预留主进程与采样线程的资源。仓库脚本中train_best.py用ray.init(num_cpus=6)+num_env_runners=4,train_ray_long.py用num_cpus=8+num_env_runners=4。

7.2 GPU 训练

config = ( PPOConfig() ... .resources(num_gpus=1) # 策略网络使用 1 张 GPU )

num_gpus控制学习器(Learner/Driver)使用 GPU 的数量;若显存充足,可配合num_gpus_per_env_runner让采样 worker 也使用 GPU。

7.3 集群训练

# 连接 Ray 集群 ray.init(address="auto") config = ( PPOConfig() ... .env_runners(num_env_runners=32) # worker 分布在集群各节点 )

address="auto"会自动发现已启动的 Ray 集群(典型启动方式为ray start --head与ray start --address=<head_ip>:6379),随后 worker 会按集群资源自动调度。需要在多机环境下,确认各节点 Python 环境与 TensorTrade 包版本一致,并确保数据 CSV 对每个 worker 节点可访问(或通过ray.put分发)。

八、模型保存、加载与导出

8.1 保存检查点

# 保存检查点,返回实际路径 checkpoint_path = algo.save('/tmp/checkpoints') print(f"Saved 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(config=config) 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_size=1000, # 调小批量 sgd_minibatch_size=128, ) .env_runners(num_env_runners=2) # 减少 worker

9.2 训练缓慢

.env_runners(num_env_runners=8) # 增加并行 worker .resources(num_gpus=1) # 或启用 GPU

9.3 训练出现 NaN

.training(lr=1e-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.4~0.9 之间(train_best.py用 0.32,train_ray_long.py用 0.5,run_ray_simulation.py用 0.9),训练初期可适当放宽。

十、替换 PPO:DQN、A2C 与 SAC

TensorTrade 默认环境的动作空间由BSH(买/卖/持有)产生,为离散动作,因此连续动作算法(SAC)需要自定义动作方案才能直接使用。

10.1 DQN(离散动作)

from ray.rllib.algorithms.dqn import DQNConfig config = ( DQNConfig() .environment(env="TradingEnv", env_config=env_config) .training( lr=1e-4, gamma=0.99, replay_buffer_config={"capacity": 100000}, ) )

10.2 A2C(比 PPO 更简单的 on-policy 算法)

from ray.rllib.algorithms.a3c import A2CConfig config = ( A2CConfig() .environment(env="TradingEnv", env_config=env_config) .training( lr=5e-4, gamma=0.99, ) )

10.3 SAC(连续动作)

from ray.rllib.algorithms.sac import SACConfig # 需要连续动作空间(需替换 BSH 动作方案) config = ( SACConfig() .environment(env="TradingEnv", env_config=env_config) .training( lr=3e-4, gamma=0.99, ) )

仓库还提供了 RLlib 的 LSTM / AttentionNet 网络集成示例(examples/use_lstm_rllib.ipynb、examples/use_attentionnet_rllib.ipynb),以及在test_ray_training.py中验证过的模型配置,可在需要时序建模时参考。

十一、要点总结与进阶路径

回顾本文核心结论:

  1. 分布式由 RLlib 托管:只需设置num_env_runners,采样、聚合、更新自动并行化;
  2. 环境工厂是成败关键:必须每次调用都创建全新环境,env.portfolio = portfolio是回调读取净值的前提;
  3. 回调承载业务指标:P&L、盈亏百分比、交易次数等全部通过episode.custom_metrics上报;
  4. PPO 是默认起点:与 TensorTrade 离散动作(BSH)天然匹配,默认超参数即来自 Optuna 调优;
  5. 务必使用独立验证集:训练 reward 上升而验证盈亏回落,就是过拟合信号。

下一步可进入 04-training/03-optuna.md,结合 examples/training/train_optuna.py 了解如何用 Optuna 自动搜索超参数;完整的实验结果汇总见 docs/EXPERIMENTS.md。

  • 人工智能
  • 金融科技
  • 机器学习

【免费下载链接】tensortrade

An open source reinforcement learning framework for training, evaluating, and deploying robust trading agents.

项目地址:https://gitcode.com/gh_mirrors/te/tensortrade
点击查看免费下载

相关推荐

上一篇:Citra模拟器终极解决方案:5步快速修复常见问题指南
下一篇:告别卡顿与模糊:GLFW视频模式让显示器性能全开

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询