13.生产部署与最佳实践
从 Jupyter Notebook 里的实验代码到真正跑在生产环境里的交易系统,这中间隔着一道不小的鸿沟。训练好的模型躺在硬盘上只是个开始,怎么让它稳定、可靠、安全地对接真实市场,才是考验功力的环节。这一章我们聊聊把 FinRL 模型部署到生产环境的关键步骤和踩过的坑。
训练模型导出与序列化
训练完成后,模型默认以 PyTorch 的 state_dict 形式存在内存里。这种格式适合继续训练或调参,但直接用于生产不够方便。我们需要把它转换成更稳定、更通用的格式。
模型格式选择
PyTorch 提供了几种序列化方式。最常用的是 torch.save 保存整个模型或仅保存权重。对于生产环境,建议只保存权重,代码里再定义网络结构。这样做的好处是模型结构可以灵活调整,不受存档文件限制。
# 保存模型权重
torch.save(agent.actor.state_dict(), 'trading_policy.pth')
# 加载时需要先实例化相同结构的网络
actor = ActorNet(state_dim, action_dim)
actor.load_state_dict(torch.load('trading_policy.pth'))
actor.eval() # 切换到评估模式,关闭 dropout 和 batch norm 的训练行为
保存权重时,记得调用 eval() 模式。这很关键,它告诉 PyTorch 关闭训练时特有的随机性,比如 dropout 的神经元失活和 batch normalization 的动量更新。生产环境需要确定性输出,忘了这行代码会导致每次预测结果不一致,调试起来让人抓狂。
版本管理与元数据
模型文件本身不记录训练时的超参数、数据版本、性能指标。这些信息在生产环境排查问题时至关重要。建议把元数据单独存成一个 JSON 文件,和模型文件同名。
import json
metadata = {
"model_type": "PPO",
"train_start_date": "2020-01-01",
"train_end_date": "2022-12-31",
"tickers": ["AAPL", "GOOGL", "MSFT"],
"state_dim": 181,
"action_dim": 3,
"sharpe_ratio": 1.85,
"max_drawdown": 0.12,
"git_commit": "a3f9b2c"
}
with open('trading_policy_meta.json', 'w') as f:
json.dump(metadata, f, indent=2)
元数据里最好包含 Git commit hash。几个月后回头看,能精确知道代码版本。tickers 列表和日期范围帮助确认模型适用的市场范围。sharpe_ratio 和 max_drawdown 是模型性能的快照,部署后监控发现指标异常下跌时,能快速判断是模型退化还是数据问题。
ONNX 格式导出
如果生产环境不用 Python,或者需要跨语言调用,ONNX 是更好的选择。它把 PyTorch 模型转换成通用计算图,支持 C++、Java、C# 等多种语言。
import torch.onnx
dummy_input = torch.randn(1, state_dim)
torch.onnx.export(
actor,
dummy_input,
"trading_policy.onnx",
input_names=['state'],
output_names=['action'],
dynamic_axes={'state': {0: 'batch_size'}, 'action': {0: 'batch_size'}},
opset_version=11
)
dummy_input 定义了输入张量的形状。dynamic_axes 参数让模型支持批量推理,生产环境往往需要一次性处理多个账户或策略。opset_version 指定算子集版本,11 是比较稳定的版本,兼容性不错。导出后,用 ONNX Runtime 加载推理,延迟比 PyTorch 低不少。
实时市场数据流接入
训练时我们用 CSV 文件或 Yahoo Finance 的静态数据。生产环境必须接入实时行情,数据延迟直接影响策略表现。
WebSocket 连接管理
大多数券商和数据提供商都提供 WebSocket 接口推送实时数据。相比轮询 REST API,WebSocket 延迟更低,也更节省资源。但连接稳定性是个挑战,网络抖动、服务器重启都会导致断线。
import websocket
import json
import threading
import time
class MarketDataStream:
def __init__(self, api_key, tickers):
self.api_key = api_key
self.tickers = tickers
self.ws = None
self.data_cache = {}
self.running = False
def on_message(self, ws, message):
data = json.loads(message)
# 缓存最新数据
self.data_cache[data['symbol']] = data
def on_error(self, ws, error):
print(f"WebSocket error: {error}")
def on_close(self, ws, close_status_code, close_msg):
print("WebSocket closed, reconnecting...")
time.sleep(5)
self.connect() # 自动重连
def on_open(self, ws):
print("WebSocket connected")
# 订阅行情
subscribe_msg = {
"action": "subscribe",
"quotes": self.tickers
}
ws.send(json.dumps(subscribe_msg))
def connect(self):
self.ws = websocket.WebSocketApp(
"wss://data.alpaca.markets/stream",
on_open=self.on_open,
on_message=self.on_message,
on_error=self.on_error,
on_close=self.on_close,
header={"APCA-API-KEY-ID": self.api_key}
)
self.running = True
self.thread = threading.Thread(target=self.ws.run_forever)
self.thread.start()
def get_latest(self, ticker):
return self.data_cache.get(ticker)
def stop(self):
self.running = False
self.ws.close()
这个类封装了 WebSocket 的生命周期管理。on_close 回调里实现自动重连,避免手动干预。data_cache 用字典缓存最新行情,策略线程调用 get_latest 时不用等待网络 IO。注意 run_forever 会阻塞当前线程,所以用 threading 在后台运行。
数据格式转换
实时推送的数据格式通常和训练时的特征向量不一致。需要写一个转换层,把原始行情加工成模型输入。
import numpy as np
import pandas as pd
class FeatureEngine:
def __init__(self, lookback_days=30):
self.lookback_days = lookback_days
self.price_history = {}
def update(self, symbol, price_data):
if symbol not in self.price_history:
self.price_history[symbol] = []
self.price_history[symbol].append(price_data)
# 保持固定长度
if len(self.price_history[symbol]) > self.lookback_days:
self.price_history[symbol].pop(0)
def get_state(self, symbol, position=0, cash=100000):
if symbol not in self.price_history or len(self.price_history[symbol]) < self.lookback_days:
return None
df = pd.DataFrame(self.price_history[symbol])
# 计算技术指标
df['sma_5'] = df['close'].rolling(5).mean()
df['sma_20'] = df['close'].rolling(20).mean()
df['rsi'] = self.calculate_rsi(df['close'])
# 取最新值
latest = df.iloc[-1]
# 构建状态向量
state = np.array([
latest['close'],
latest['sma_5'],
latest['sma_20'],
latest['rsi'],
position,
cash
])
return state
def calculate_rsi(self, prices, window=14):
delta = prices.diff()
gain = (delta.where(delta > 0, 0)).rolling(window=window).mean()
loss = (-delta.where(delta < 0, 0)).rolling(window=window).mean()
rs = gain / loss
rsi = 100 - (100 / (1 + rs))
return rsi
FeatureEngine 维护每个标的的价格历史,计算 SMA 和 RSI 等技术指标。get_state 方法返回模型需要的完整状态向量,包括持仓和现金信息。这里有个细节:历史数据不足时返回 None,策略层需要处理这种情况,避免用不完整数据做决策。
缓存与降级
实时数据源难免出问题,必须有降级方案。可以缓存最近一小时的数据,当实时流断开时,用缓存数据继续运行,但降低交易频率或暂停新开仓。
from datetime import datetime, timedelta
class DataManager:
def __init__(self, stream, feature_engine):
self.stream = stream
self.feature_engine = feature_engine
self.last_update = {}
self.fallback_mode = False
def get_state_with_fallback(self, symbol):
# 检查数据新鲜度
now = datetime.now()
if symbol in self.last_update:
time_diff = now - self.last_update[symbol]
if time_diff > timedelta(seconds=30):
self.fallback_mode = True
print(f"Data stale for {symbol}, entering fallback mode")
# 尝试获取实时数据
raw_data = self.stream.get_latest(symbol)
if raw_data:
self.feature_engine.update(symbol, raw_data)
self.last_update[symbol] = now
self.fallback_mode = False
# 获取特征
state = self.feature_engine.get_state(symbol)
# 降级处理
if state is None and self.fallback_mode:
print("Using cached data for decision")
# 返回上一次的 state,或保守的默认状态
return self.get_last_known_state(symbol)
return state
这个管理器增加了数据新鲜度检查。30 秒没更新就进入降级模式,避免用过时数据交易。fallback_mode 可以触发策略层的保守逻辑,比如只平仓不开新仓。实际部署时,降级阈值要根据标的的流动性调整,外汇和加密货币可以容忍更长的延迟。
模拟交易与实盘对接
模拟交易(Paper Trading)是连接回测和实盘的桥梁。它用真实市场数据执行虚拟订单,能验证策略在真实滑点和佣金下的表现,同时避免真金白银的风险。
订单执行逻辑
FinRL 的模拟交易模块需要对接券商 API。这里以 Alpaca 为例,展示订单管理的核心逻辑。
import alpaca_trade_api as tradeapi
class PaperTrader:
def __init__(self, api_key, secret_key, base_url='https://paper-api.alpaca.markets'):
self.api = tradeapi.REST(api_key, secret_key, base_url, api_version='v2')
self.positions = {}
self.order_history = []
def place_order(self, symbol, side, qty, order_type='market', time_in_force='day'):
try:
order = self.api.submit_order(
symbol=symbol,
side=side,
qty=qty,
type=order_type,
time_in_force=time_in_force
)
self.order_history.append({
'symbol': symbol,
'side': side,
'qty': qty,
'status': order.status,
'submitted_at': order.submitted_at
})
print(f"Order submitted: {side} {qty} shares of {symbol}")
return order
except Exception as e:
print(f"Order failed: {e}")
return None
def get_position(self, symbol):
try:
position = self.api.get_position(symbol)
return {
'symbol': symbol,
'qty': int(position.qty),
'avg_entry_price': float(position.avg_entry_price),
'market_value': float(position.market_value)
}
except:
return None # 无持仓
def cancel_all_orders(self):
self.api.cancel_all_orders()
print("All open orders cancelled")
PaperTrader 封装了下单、查持仓、撤单等基本操作。submit_order 会返回订单对象,包含状态和执行价格。生产环境必须捕获异常,网络超时或资金不足都会导致下单失败。order_history 记录所有操作,方便事后审计和性能分析。
策略循环集成
把模型预测和交易执行串起来,形成完整的策略循环。
import time
from datetime import datetime
class TradingLoop:
def __init__(self, trader, data_manager, model, tickers):
self.trader = trader
self.data_manager = data_manager
self.model = model
self.tickers = tickers
self.is_running = False
def run(self):
self.is_running = True
# 等待市场开盘
self.wait_for_market_open()
while self.is_running:
try:
# 检查是否收盘
if not self.is_market_open():
print("Market closed, stopping trading loop")
break
for ticker in self.tickers:
self.trade_single_ticker(ticker)
# 等待下一个决策周期
time.sleep(60) # 每分钟决策一次
except Exception as e:
print(f"Error in trading loop: {e}")
time.sleep(5) # 出错后短暂休眠再试
def trade_single_ticker(self, ticker):
# 获取状态
position = self.trader.get_position(ticker)
current_pos = position['qty'] if position else 0
state = self.data_manager.get_state_with_fallback(ticker, position=current_pos)
if state is None:
return # 数据不足,跳过本次决策
# 模型预测
with torch.no_grad():
action_probs = self.model(torch.FloatTensor(state).unsqueeze(0))
action = torch.argmax(action_probs, dim=1).item()
# 执行动作
if action == 1: # 买入
self.execute_buy(ticker)
elif action == 2: # 卖出
self.execute_sell(ticker)
else: # 持有
pass
def execute_buy(self, ticker):
# 检查购买力
account = self.trader.api.get_account()
buying_power = float(account.buying_power)
# 计算购买数量
last_price = self.data_manager.get_latest_price(ticker)
if not last_price:
return
qty = int(buying_power * 0.1 / last_price) # 用 10% 资金
if qty > 0:
self.trader.place_order(ticker, 'buy', qty)
def execute_sell(self, ticker):
position = self.trader.get_position(ticker)
if position and position['qty'] > 0:
self.trader.place_order(ticker, 'sell', position['qty'])
def stop(self):
self.is_running = False
self.trader.cancel_all_orders()
TradingLoop 是策略的核心驱动。wait_for_market_open 确保只在交易时段运行,避免盘前盘后流动性不足的风险。trade_single_ticker 里,模型输出动作后,执行层要检查资金和持仓,避免超买或裸卖空。这里只用了 10% 资金买入,是风险控制的一种简单实现。实际部署时,这个比例应该根据策略波动率和账户风险承受能力动态调整。
风险管理与限额
生产环境必须设置硬性的风险限额,防止模型失控。
class RiskManager:
def __init__(self, max_position_per_stock=100, max_daily_loss=0.05):
self.max_position_per_stock = max_position_per_stock
self.max_daily_loss = max_daily_loss
self.daily_pnl = 0
self.last_equity = None
def check_position_limit(self, symbol, proposed_qty):
"""检查持仓限额"""
return abs(proposed_qty) <= self.max_position_per_stock
def check_daily_loss_limit(self, current_equity):
"""检查当日亏损限额"""
if self.last_equity is None:
self.last_equity = current_equity
self.daily_pnl = (current_equity - self.last_equity) / self.last_equity
if self.daily_pnl < -self.max_daily_loss:
print(f"Daily loss limit reached: {self.daily_pnl:.2%}")
return False
return True
def reset_daily_limit(self, opening_equity):
"""每个交易日重置"""
self.last_equity = opening_equity
self.daily_pnl = 0
RiskManager 提供两层保护。持仓限额防止单个标的风险过度集中。每日亏损限额是最后一道防线,当亏损达到 5% 时,策略应该停止交易。这个类需要在 TradingLoop 的每个决策周期调用,一旦触发限制,立刻暂停策略并发送告警。
系统监控与异常告警
模型上线后,需要持续监控它的健康状况。监控分三个层面:数据质量、模型性能、系统资源。
日志记录最佳实践
打印到控制台在 Notebook 里够用,生产环境需要持久化日志。
import logging
from logging.handlers import RotatingFileHandler
def setup_logging():
logger = logging.getLogger('trading_system')
logger.setLevel(logging.INFO)
# 文件处理器,按大小轮转
file_handler = RotatingFileHandler(
'trading.log',
maxBytes=10*1024*1024, # 10MB
backupCount=5
)
# 控制台处理器
console_handler = logging.StreamHandler()
# 格式化
formatter = logging.Formatter(
'%(asctime)s - %(name)s - %(levelname)s - %(message)s'
)
file_handler.setFormatter(formatter)
console_handler.setFormatter(formatter)
logger.addHandler(file_handler)
logger.addHandler(console_handler)
return logger
# 使用
logger = setup_logging()
logger.info("Trading loop started")
logger.warning("Data stale, using fallback mode")
logger.error("Order execution failed", exc_info=True)
RotatingFileHandler 自动管理日志文件大小,超过 10MB 就创建新文件,保留最近 5 个。这样避免日志占满磁盘。exc_info=True 会把异常堆栈一起记录,排查问题时能精确定位。日志要分级,INFO 记录正常操作,WARNING 记录降级行为,ERROR 记录失败操作。
性能指标收集
监控策略的实时表现,及时发现性能退化。
import sqlite3
from datetime import datetime
class PerformanceTracker:
def __init__(self, db_path='trading_metrics.db'):
self.conn = sqlite3.connect(db_path)
self._create_table()
def _create_table(self):
cursor = self.conn.cursor()
cursor.execute('''
CREATE TABLE IF NOT EXISTS metrics (
timestamp TEXT,
symbol TEXT,
action INTEGER,
position INTEGER,
pnl REAL,
portfolio_value REAL
)
''')
self.conn.commit()
def log_decision(self, symbol, action, position, pnl, portfolio_value):
cursor = self.conn.cursor()
cursor.execute('''
INSERT INTO metrics VALUES (?, ?, ?, ?, ?, ?)
''', (
datetime.now().isoformat(),
symbol,
action,
position,
pnl,
portfolio_value
))
self.conn.commit()
def get_daily_sharpe(self):
cursor = self.conn.cursor()
cursor.execute('''
SELECT pnl FROM metrics
WHERE timestamp >= date('now', '-1 day')
''')
returns = [row[0] for row in cursor.fetchall()]
if len(returns) < 2:
return 0
return np.mean(returns) / np.std(returns) * np.sqrt(252)
用 SQLite 存储每个决策周期的指标,轻量且不需要额外服务。get_daily_sharpe 计算实时夏普比率,和回测时的基准对比,如果连续几天都低于预期,说明模型可能过拟合或市场 regime 已经改变。
异常告警通知
关键异常需要立即通知,不能靠人盯着日志。
import smtplib
from email.mime.text import MIMEText
class AlertManager:
def __init__(self, smtp_config):
self.smtp_host = smtp_config['host']
self.smtp_port = smtp_config['port']
self.username = smtp_config['username']
self.password = smtp_config['password']
self.recipients = smtp_config['recipients']
def send_alert(self, subject, body, level='WARNING'):
if level == 'INFO':
return # 只发送警告和错误
msg = MIMEText(body)
msg['Subject'] = f"[Trading System] {level}: {subject}"
msg['From'] = self.username
msg['To'] = ', '.join(self.recipients)
try:
server = smtplib.SMTP(self.smtp_host, self.smtp_port)
server.starttls()
server.login(self.username, self.password)
server.send_message(msg)
server.quit()
print("Alert sent successfully")
except Exception as e:
print(f"Failed to send alert: {e}")
def check_system_health(self, data_manager, trader):
"""检查系统健康状态"""
alerts = []
# 检查数据延迟
if data_manager.fallback_mode:
alerts.append("Market data is in fallback mode")
# 检查持仓异常
positions = trader.api.list_positions()
total_value = sum(float(p.market_value) for p in positions)
if total_value > 1000000: # 假设限额 100 万
alerts.append(f"Position limit exceeded: ${total_value:,.2f}")
# 发送汇总告警
if alerts:
self.send_alert(
"System Health Check Failed",
"\n".join(alerts),
level='ERROR'
)
AlertManager 在关键异常时发邮件。check_system_health 定期执行,检查数据降级和持仓限额。生产环境建议接入专业的监控服务,如 Prometheus + Grafana 或 Datadog,它们提供更强大的告警规则和可视化。但邮件告警作为保底手段依然必要,简单可靠。
监控面板搭建
命令行输出不够直观,可以搭个简单的 Web 面板查看实时状态。
from flask import Flask, jsonify
import threading
app = Flask(__name__)
class Dashboard:
def __init__(self, trader, performance_tracker):
self.trader = trader
self.tracker = performance_tracker
def get_status(self):
account = self.trader.api.get_account()
positions = self.trader.api.list_positions()
return {
'equity': float(account.equity),
'cash': float(account.cash),
'buying_power': float(account.buying_power),
'position_count': len(positions),
'daily_sharpe': self.tracker.get_daily_sharpe()
}
dashboard = Dashboard(trader, performance_tracker)
@app.route('/api/status')
def status():
return jsonify(dashboard.get_status())
def run_dashboard():
app.run(host='0.0.0.0', port=5000, debug=False)
# 在后台启动
dashboard_thread = threading.Thread(target=run_dashboard)
dashboard_thread.start()
Flask 提供 REST API,前端用简单 HTML 或 Grafana 都能对接。get_status 返回账户概况和实时夏普比率。监控面板不要和交易循环跑在同一个线程,避免阻塞。生产环境建议用 uWSGI 或 Gunicorn 部署 Flask,比内置服务器更稳定。
部署架构建议
把上面模块串起来,一个基础的生产架构大概长这样:
数据流层: MarketDataStream -> FeatureEngine -> DataManager
决策层: TradingLoop -> Model -> RiskManager
执行层: PaperTrader -> Broker API
监控层: PerformanceTracker -> AlertManager -> Dashboard
每层之间通过队列或缓存解耦,避免单点故障拖垮整个系统。比如 DataManager 把状态写入 Redis,TradingLoop 从 Redis 读取,这样即使数据流重启,决策层也能继续运行。
部署时建议用 Docker 容器化,每个服务独立运行。Alpaca 的 paper trading 环境是免费的,适合验证整个流程。等模拟交易跑稳了,再换到实盘 API,只需改 base_url 和 API Key。
模型更新策略也很重要。可以每周日用新数据重新训练,生成新模型。部署时采用蓝绿部署:先启动新模型实例,逐步切流量,监控一小时无异常再下线旧模型。这样能最大限度减少服务中断。
生产部署没有银弹,核心在于监控和回滚机制。每个环节都要假设会失败,数据会断、模型会错、网络会抖。只有把这些失败场景都考虑到,系统才能在真实市场的高压下存活。
下一章我们会深入讨论当系统真的出问题时的排查思路和性能优化技巧,包括怎么定位数据维度不匹配、训练不稳定这些常见坑。