- 人工智能
- 金融科技
- 机器学习
【免费下载链接】tensortrade
An open source reinforcement learning framework for training, evaluating, and deploying robust trading agents.
导读
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) # 减少 worker9.2 训练缓慢
.env_runners(num_env_runners=8) # 增加并行 worker .resources(num_gpus=1) # 或启用 GPU9.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中验证过的模型配置,可在需要时序建模时参考。
十一、要点总结与进阶路径
回顾本文核心结论:
- 分布式由 RLlib 托管:只需设置
num_env_runners,采样、聚合、更新自动并行化; - 环境工厂是成败关键:必须每次调用都创建全新环境,
env.portfolio = portfolio是回调读取净值的前提; - 回调承载业务指标:P&L、盈亏百分比、交易次数等全部通过
episode.custom_metrics上报; - PPO 是默认起点:与 TensorTrade 离散动作(BSH)天然匹配,默认超参数即来自 Optuna 调优;
- 务必使用独立验证集:训练 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.
相关推荐
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 为主线,系统讲解该
人工智能金融科技机器学习
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考