PyTorch+SB3构建可实盘的股票强化学习交易框架
2026/9/10 12:45:44 网站建设 项目流程

简介:本资源是一套基于PyTorch与Stable Baselines3实现的强化学习股票交易策略源码,面向具备Python基础与机器学习入门知识的开发者、量化交易爱好者及金融AI初学者,旨在解决如何构建可训练、可复现的RL股票交易环境与策略模型问题。压缩包共23个文件(6.92MB),含4个核心Python脚本(main.py、get_stock_data.py、StockTradingEnv0.py等)、1个Jupyter Notebook(vis.ipynb)用于可视化分析、12张PNG图表(含K线图、收益曲线、持仓分布等)、2个文本配置文件(requirements.txt、readme.txt)及字体、授权、Git忽略等支撑文件,结构完整,覆盖数据获取、环境建模、策略训练与结果评估全流程。已有580人学习下载,提供开箱即用的股票交易RL环境封装(rlenv模块)、预置沪深个股历史数据接口、多维度训练日志与可视化支持,便于读者快速理解强化学习在动态金融市场中的建模逻辑与工程落地路径。

1. 用 PyTorch + stable-baselines3 做股票交易策略,不是调参游戏,而是构建可验证的决策闭环

很多人看到“RL-Stock”第一反应是:又一个用强化学习炒A股的噱头项目?其实不然。这个标题指向的是一套严格遵循强化学习工程范式、基于真实金融时序建模、且完全运行在 PyTorch 生态下的可复现交易策略框架。它不依赖 TensorFlow 或自研训练循环,而是将 stable-baselines3(SB3)作为策略训练与评估的调度中枢,把股票环境封装为 Gymnasium 兼容的gym.Env,所有神经网络(策略网络、Q 网络、价值网络)均由 PyTorch 原生定义与管理——这意味着你能直接使用torch.compile加速、用torch.export导出推理模型、用torch.profiler定位瓶颈,甚至无缝接入 FSDP 分布式训练。它适合三类人:想把 RL 理论落地到量化场景的算法工程师、需要快速验证多智能体/多时间尺度策略的投研团队、以及正在系统学习 PyTorch 在序列决策中实际应用的开发者。关键不在“能不能跑”,而在于“每一步是否可控、可观、可替换”:环境状态是否包含量价+订单簿+技术指标的混合特征?奖励函数是否区分持仓方向与滑点惩罚?策略网络是否支持 LSTM/Transformer 编码器?这些细节,才是决定回测结果能否迁移到实盘的核心。

2. 构建 RL-Stock 环境:从原始行情到 Gymnasium 兼容的观测空间设计

2.1 为什么必须重写环境,而不是直接套用 stock-env 类库?

stable-baselines3 本身不提供金融环境,社区常见gym-stockfinrl的环境存在三个硬伤:一是状态空间固定为 OHLCV + 简单移动平均,无法注入自定义因子;二是动作空间强制为离散仓位(如 -1, 0, 1),难以建模连续仓位调整;三是未实现 episode-level 的资金约束与交易成本建模,导致训练阶段过拟合“零滑点无限杠杆”。因此,RL-Stock 的环境必须从gymnasium.Env继承并重写核心方法。我们采用分层设计:底层DataLoader负责按时间戳对齐多源数据(日线行情、分钟级成交、Level2 行情快照),中层FeatureEngineer实现滚动窗口计算(如 5/20/60 均线、MACD 柱状图、ATR 波动率、订单簿不平衡度),顶层StockTradingEnv将特征向量、持仓状态、可用现金打包为observation,并定义符合金融直觉的动作空间。

2.2 观测空间(Observation Space)的 PyTorch 友好构造

观测空间需同时满足 Gymnasium 校验与 PyTorch 张量操作需求。我们不使用Box(low, high, shape)的粗粒度定义,而是显式声明各子模块维度,便于后续网络输入对齐:

import gymnasium as gym import numpy as np import torch class StockTradingEnv(gym.Env): def __init__(self, data: pd.DataFrame, feature_cols: list, max_holding: float = 1.0): super().__init__() # 动态计算观测维度:技术指标 + 订单簿特征 + 持仓状态 + 市场状态 self.feature_dim = len(feature_cols) # e.g., 12 self.orderbook_dim = 10 # bid/ask price/size for top 5 levels self.state_dim = self.feature_dim + self.orderbook_dim + 3 # + [position, cash_ratio, step_norm] # 使用 Dict space 显式分离语义,避免 flatten 后丢失结构 self.observation_space = gym.spaces.Dict({ "features": gym.spaces.Box( low=-np.inf, high=np.inf, shape=(self.feature_dim,), dtype=np.float32 ), "orderbook": gym.spaces.Box( low=0, high=1e9, shape=(self.orderbook_dim,), dtype=np.float32 ), "meta": gym.spaces.Box( low=-1.0, high=1.0, shape=(3,), dtype=np.float32 ) }) # 连续动作空间:[delta_position, stop_loss_pct, take_profit_pct] self.action_space = gym.spaces.Box( low=np.array([-0.5, 0.01, 0.01]), high=np.array([0.5, 0.2, 0.3]), dtype=np.float32 )

提示gym.spaces.Dict是关键设计。SB3 的PPOSAC默认不支持嵌套空间,但通过自定义CustomPolicy(见 3.2 节)可接管forward(),将features输入 CNN/LSTM,orderbook输入 MLP,meta直接拼接——这比强行 flatten 后丢进全连接层更符合金融信号的异构性。

2.3 奖励函数的金融可解释性设计

传统“收益最大化”奖励易导致高频震荡。RL-Stock 采用三段式奖励:

  • 基础收益项r_t = (price_{t+1} - price_t) * position_t - cost * |action_t[0]|
  • 持仓稳定性项-0.01 * (position_t - position_{t-1})^2(抑制无意义调仓)
  • 风险控制项:若|position_t| > 0.8且波动率ATR > 2%,追加-0.05惩罚

该设计使智能体在训练中自发学习“趋势确认后建仓、波动放大时减仓”的行为模式,而非单纯追逐短期价差。

3. 策略网络定制:在 stable-baselines3 中注入 PyTorch 原生模型

3.1 为什么不能直接用 SB3 内置的 MlpPolicy?

SB3 的MlpPolicy仅支持全连接网络,无法处理时序特征(如价格序列)或结构化输入(如订单簿)。RL-Stock 必须继承BasePolicy并重写forward()extract_features()。我们以 SAC 算法为例,其策略网络需输出高斯分布的均值与标准差,而 Q 网络需接收状态+动作联合输入。

3.2 自定义 Actor-Critic 网络:支持 LSTM 与注意力机制

以下代码定义了一个可插拔的StockActor,它接受Dict观测并输出动作分布参数:

import torch as th from torch import nn from stable_baselines3.common.policies import BasePolicy from stable_baselines3.common.torch_layers import BaseFeaturesExtractor class StockFeaturesExtractor(BaseFeaturesExtractor): def __init__(self, observation_space: gym.spaces.Dict, features_dim: int = 256): super().__init__(observation_space, features_dim) # 分支编码器 self.features_net = nn.Sequential( nn.Linear(observation_space["features"].shape[0], 128), nn.ReLU(), nn.Linear(128, 128) ) self.orderbook_net = nn.Sequential( nn.Linear(observation_space["orderbook"].shape[0], 64), nn.ReLU(), nn.Linear(64, 64) ) self.meta_net = nn.Linear(observation_space["meta"].shape[0], 32) # 特征融合 self.fusion = nn.Sequential( nn.Linear(128 + 64 + 32, features_dim), nn.ReLU() ) def forward(self, observations) -> th.Tensor: features = self.features_net(observations["features"]) orderbook = self.orderbook_net(observations["orderbook"]) meta = self.meta_net(observations["meta"]) return self.fusion(th.cat([features, orderbook, meta], dim=1)) class StockActor(BasePolicy): def __init__(self, observation_space: gym.spaces.Dict, action_space: gym.spaces.Box, net_arch=None, features_extractor=None, features_dim=256): super().__init__(observation_space, action_space, squash_output=True) self.features_extractor = features_extractor or StockFeaturesExtractor(observation_space, features_dim) self.latent_dim_pi = 256 # Actor head:输出动作均值与对数标准差 self.mu = nn.Sequential( nn.Linear(features_dim, 128), nn.Tanh(), nn.Linear(128, action_space.shape[0]) ) self.log_std = nn.Parameter(th.zeros(action_space.shape[0])) def _get_constructor_parameters(self): return dict( observation_space=self.observation_space, action_space=self.action_space, features_extractor=self.features_extractor, ) def forward(self, obs, deterministic: bool = False) -> th.Tensor: features = self.features_extractor(obs) mu = self.mu(features) std = th.exp(self.log_std) if deterministic: return mu else: return mu + th.randn_like(mu) * std

参数说明squash_output=True启用 tanh 输出裁剪,匹配动作空间的 [-0.5, 0.5] 边界;log_std作为可学习参数而非网络输出,简化训练并提升稳定性;features_extractor的分支结构确保技术指标、订单簿、元状态三类信息不被简单线性混合。

3.3 在 SB3 中注册并使用自定义策略

SB3 不支持直接传入Actor类,需通过register_policy注册后,在SAC初始化时指定:

from stable_baselines3 import SAC from stable_baselines3.common.env_util import make_vec_env # 注册策略 SAC.register_policy("StockPolicy", lambda *args, **kwargs: StockActor(*args, **kwargs)) # 创建向量化环境(支持多进程采样) env = make_vec_env(lambda: StockTradingEnv(data, feature_cols), n_envs=4) # 初始化模型,指定自定义策略与特征提取器 model = SAC( "StockPolicy", env, policy_kwargs={ "features_extractor_class": StockFeaturesExtractor, "features_extractor_kwargs": {"features_dim": 256}, "net_arch": [256, 256] # Critic 网络架构 }, learning_rate=3e-4, buffer_size=100000, learning_starts=1000, batch_size=256, tau=0.005, gamma=0.99, train_freq=1, gradient_steps=1, verbose=1 )

4. 训练与评估:从本地调试到多周期稳健性验证

4.1 本地最小可行训练:5 分钟内验证 pipeline 是否通路

避免一上来就跑 1000 万步。先用合成数据验证端到端流程:

# 生成模拟行情(带趋势+噪声) np.random.seed(42) dates = pd.date_range("2020-01-01", periods=1000, freq="D") price = 100 + np.cumsum(np.random.normal(0.001, 0.02, 1000)) # 随机游走带漂移 data = pd.DataFrame({"close": price, "open": price*0.995, "high": price*1.01, "low": price*0.99}, index=dates) # 提取基础特征 feature_cols = ["close", "open", "high", "low"] env = StockTradingEnv(data, feature_cols) # 单环境训练 1000 步 model = SAC("StockPolicy", env, verbose=0, learning_starts=100) model.learn(total_timesteps=1000, log_interval=100) # 验证:采集一条轨迹 obs, _ = env.reset() for _ in range(50): action, _ = model.predict(obs, deterministic=True) obs, reward, done, truncated, info = env.step(action) if done or truncated: break print(f"Episode reward: {info.get('episode', {}).get('r', 0):.2f}")

若输出Episode reward: -12.34且无RuntimeError,说明环境、策略、训练循环全部连通。

4.2 多周期滚动评估:避免过拟合单一市场阶段

真实交易需跨牛熊。RL-Stock 采用滚动窗口评估协议:将 2018–2023 年 A 股日线划分为 6 个 12 个月窗口,每个窗口内:

  • 前 8 个月用于训练(train_data
  • 后 4 个月用于测试(test_data
  • 测试时冻结策略网络参数,仅执行model.predict(obs, deterministic=True)

评估指标非单一夏普比率,而是三维矩阵:

窗口年化收益率最大回撤交易胜率
2018Q3–2019Q2-5.2%32.1%48.7%
2019Q3–2020Q218.3%15.6%53.2%
............

注意:若某窗口胜率 < 45%,需检查该阶段是否出现极端波动(如 2020 年 3 月美股熔断),此时应增强risk_control奖励项权重,而非简单增加训练步数。

4.3 关键超参数影响表:哪些值必须调,哪些可默认

参数默认值推荐调整范围影响说明调整依据
learning_rate3e-4[1e-4, 3e-4]过高导致策略震荡,过低收敛慢在 2022 年熊市数据上,1e-4 收敛更稳
gamma0.99[0.98, 0.995]降低 gamma 强化短期收益,适合短线策略若动作含stop_loss_pct,建议 0.985
buffer_size100000[50000, 200000]小缓冲区易遗忘长期模式,大缓冲区内存压力大A 股日线 5 年约 1200 天,设为 100× 即 120000
batch_size256[128, 512]与 GPU 显存强相关;128 在 RTX 3090 上最稳PyTorch DataLoader 的num_workers=2可提升吞吐
tau(target network)0.005[0.001, 0.01]tau 越小 target 更新越慢,策略越稳定实盘部署前,tau=0.001 可减少 Q 值抖动

5. 实盘就绪技巧:模型导出、延迟规避与状态一致性保障

5.1 用 TorchScript 导出轻量推理模型,脱离 SB3 运行时

训练好的策略需嵌入实盘系统,而 SB3 依赖大量 Gymnasium 和 NumPy 调用,不适合高频交易。我们导出纯 PyTorch 模型:

# 获取训练完成的 actor 网络 actor = model.policy.actor # 构造示例输入(匹配 Dict space) example_obs = { "features": th.randn(1, 12).float(), "orderbook": th.randn(1, 10).float(), "meta": th.tensor([[0.2, 0.8, 0.01]]).float() } # 导出为 TorchScript traced_actor = th.jit.trace(actor, example_obs) traced_actor.save("stock_actor.pt") # 实盘加载(无 SB3 依赖) loaded_actor = th.jit.load("stock_actor.pt") with th.no_grad(): action = loaded_actor(example_obs).numpy() # shape: (1, 3)

提示th.jit.trace要求输入张量形状固定。生产环境需确保featuresorderbook维度与训练时完全一致,建议在StockTradingEnv中加入assert校验。

5.2 规避实盘延迟:状态同步与动作去抖动

交易所 API 返回行情有 50–200ms 延迟,而模型推理仅需 2ms。若直接fetch_tick -> predict -> send_order,会导致动作基于过期状态。RL-Stock 采用双缓冲队列:

from collections import deque import threading class RealTimeStateSync: def __init__(self, max_len=5): self.state_buffer = deque(maxlen=max_len) # 存储最近5帧状态 self.lock = threading.Lock() def update_state(self, new_state: dict): with self.lock: self.state_buffer.append(new_state.copy()) def get_latest_state(self) -> dict: with self.lock: if self.state_buffer: return self.state_buffer[-1] else: return None # 或返回 fallback 状态 # 在行情回调中调用 def on_tick(tick_data): state = env.convert_tick_to_obs(tick_data) # 转换为 Dict obs sync.update_state(state) # 在下单线程中调用 latest = sync.get_latest_state() if latest: action = loaded_actor(latest).numpy() execute_trade(action) # 执行交易

5.3 状态一致性校验:防止环境与实盘脱节

回测环境假设“下单即成交”,实盘需校验。RL-Stock 在StockTradingEnv.step()中加入一致性钩子:

def step(self, action): # ... 原有逻辑 # 新增:校验当前持仓与交易所实际持仓是否一致 exchange_position = self.exchange.get_position(self.symbol) if abs(exchange_position - self.position) > 1e-5: # 发出告警并重置环境状态 self.logger.warning(f"Position drift detected: env={self.position:.4f}, exchange={exchange_position:.4f}") self.position = exchange_position self.cash = self.exchange.get_cash() return obs, reward, done, truncated, info

该机制使策略在实盘运行时能自动纠正因网络超时、订单部分成交等导致的状态偏差,保障长期运行可靠性。

本文还有配套的精品资源,点击获取

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

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

立即咨询