1. Deployment Architecture
Deploying reinforcement learning trading systems in production requires careful consideration of system architecture, real-time execution, monitoring, and risk management to ensure reliable and profitable operation.
Deployment Components:
- Model Serving: Real-time inference with low latency
- Data Pipeline: Real-time market data processing
- Execution Engine: Order management and execution
- Risk Management: Real-time risk monitoring and controls
- Monitoring: System health and performance tracking
System Architecture
# Production RL Trading System Architecture
class ProductionTradingSystem:
def __init__(self, config):
self.config = config
# Core components
self.data_pipeline = DataPipeline(config['data'])
self.model_server = ModelServer(config['model'])
self.execution_engine = ExecutionEngine(config['execution'])
self.risk_manager = RiskManager(config['risk'])
self.monitoring = MonitoringSystem(config['monitoring'])
# State management
self.portfolio_state = PortfolioState()
self.market_state = MarketState()
def start_trading(self):
"""Start the trading system"""
# Initialize components
self.data_pipeline.start()
self.model_server.start()
self.execution_engine.start()
self.risk_manager.start()
self.monitoring.start()
# Start main trading loop
self._trading_loop()
def _trading_loop(self):
"""Main trading loop"""
while self.running:
try:
# Get latest market data
market_data = self.data_pipeline.get_latest_data()
# Update market state
self.market_state.update(market_data)
# Check risk limits
if not self.risk_manager.check_limits(self.portfolio_state):
continue
# Get model prediction
action = self.model_server.predict(self.market_state, self.portfolio_state)
# Validate action
if self.risk_manager.validate_action(action, self.portfolio_state):
# Execute trade
execution_result = self.execution_engine.execute(action)
# Update portfolio state
self.portfolio_state.update(execution_result)
# Log trade
self.monitoring.log_trade(execution_result)
# Update monitoring metrics
self.monitoring.update_metrics(self.portfolio_state, self.market_state)
except Exception as e:
self.monitoring.log_error(e)
self.risk_manager.handle_error(e)
time.sleep(0.1) # 100ms loop
2. Model Serving Infrastructure
Real-time Inference
import torch
import asyncio
import redis
from fastapi import FastAPI
import uvicorn
class ModelServer:
def __init__(self, config):
self.config = config
self.model = self._load_model()
self.preprocessor = self._load_preprocessor()
self.cache = redis.Redis(host=config['redis_host'], port=config['redis_port'])
# Performance tracking
self.inference_times = []
self.cache_hits = 0
self.cache_misses = 0
def _load_model(self):
"""Load trained RL model"""
model = DQNNetwork(input_size=20, hidden_sizes=[128, 64], output_size=3)
checkpoint = torch.load(self.config['model_path'], map_location='cpu')
model.load_state_dict(checkpoint['model_state_dict'])
model.eval()
return model
def predict(self, market_state, portfolio_state):
"""Get model prediction with caching"""
# Create cache key
cache_key = self._create_cache_key(market_state, portfolio_state)
# Check cache first
cached_result = self.cache.get(cache_key)
if cached_result:
self.cache_hits += 1
return json.loads(cached_result)
self.cache_misses += 1
# Preprocess input
start_time = time.time()
input_tensor = self.preprocessor.transform(market_state, portfolio_state)
# Model inference
with torch.no_grad():
q_values = self.model(input_tensor)
action = q_values.argmax().item()
confidence = torch.softmax(q_values, dim=1).max().item()
inference_time = time.time() - start_time
self.inference_times.append(inference_time)
# Cache result
result = {'action': action, 'confidence': confidence, 'q_values': q_values.tolist()}
self.cache.setex(cache_key, 60, json.dumps(result)) # 1 minute cache
return result
def _create_cache_key(self, market_state, portfolio_state):
"""Create cache key from state"""
market_hash = hash(tuple(market_state.flatten()))
portfolio_hash = hash(tuple(portfolio_state.flatten()))
return f"prediction:{market_hash}:{portfolio_hash}"
def get_performance_stats(self):
"""Get model serving performance statistics"""
if not self.inference_times:
return {'avg_inference_time': 0, 'cache_hit_rate': 0}
return {
'avg_inference_time': np.mean(self.inference_times),
'p95_inference_time': np.percentile(self.inference_times, 95),
'cache_hit_rate': self.cache_hits / (self.cache_hits + self.cache_misses),
'total_predictions': self.cache_hits + self.cache_misses
}
# FastAPI server for model serving
app = FastAPI(title="RL Trading Model Server")
model_server = None
@app.on_event("startup")
async def startup_event():
global model_server
config = load_config()
model_server = ModelServer(config)
@app.post("/predict")
async def predict(request: PredictionRequest):
"""API endpoint for model predictions"""
try:
result = model_server.predict(request.market_state, request.portfolio_state)
return {"status": "success", "result": result}
except Exception as e:
return {"status": "error", "message": str(e)}
@app.get("/health")
async def health_check():
"""Health check endpoint"""
stats = model_server.get_performance_stats()
return {"status": "healthy", "stats": stats}
Model Optimization
Production Optimization:
class ModelOptimizer:
def __init__(self, model):
self.model = model
def optimize_for_production(self):
"""Optimize model for production deployment"""
# Quantization
quantized_model = self._quantize_model()
# Pruning
pruned_model = self._prune_model(quantized_model)
# Compilation
compiled_model = self._compile_model(pruned_model)
return compiled_model
def _quantize_model(self):
"""Quantize model for faster inference"""
import torch.quantization as quantization
self.model.eval()
self.model.qconfig = quantization.get_default_qconfig('fbgemm')
# Prepare for quantization
model_prepared = quantization.prepare(self.model)
# Calibrate with sample data
calibration_data = self._get_calibration_data()
with torch.no_grad():
for data in calibration_data:
model_prepared(data)
# Convert to quantized model
quantized_model = quantization.convert(model_prepared)
return quantized_model
def _prune_model(self, model, sparsity=0.2):
"""Prune model to reduce size"""
import torch.nn.utils.prune as prune
# Prune linear layers
for module in model.modules():
if isinstance(module, torch.nn.Linear):
prune.l1_unstructured(module, name='weight', amount=sparsity)
prune.remove(module, 'weight')
return model
def _compile_model(self, model):
"""Compile model for optimal performance"""
try:
# PyTorch 2.0 compilation
compiled_model = torch.compile(model)
return compiled_model
except:
# Fallback for older versions
return model
3. Risk Management in Production
Real-time Risk Monitoring
class ProductionRiskManager:
def __init__(self, config):
self.config = config
# Risk limits
self.max_position_size = config['max_position_size']
self.max_drawdown = config['max_drawdown']
self.max_daily_loss = config['max_daily_loss']
self.var_limit = config['var_limit']
# Risk state
self.daily_pnl = 0
self.portfolio_value = config['initial_capital']
self.max_portfolio_value = config['initial_capital']
# Alerts
self.alerts = []
self.circuit_breakers = []
def check_limits(self, portfolio_state):
"""Check all risk limits"""
# Position size limit
if not self._check_position_limit(portfolio_state):
return False
# Drawdown limit
if not self._check_drawdown_limit(portfolio_state):
return False
# Daily loss limit
if not self._check_daily_loss_limit():
return False
# VaR limit
if not self._check_var_limit(portfolio_state):
return False
return True
def _check_position_limit(self, portfolio_state):
"""Check position size limits"""
position_ratio = abs(portfolio_state['position_value']) / portfolio_state['total_value']
if position_ratio > self.max_position_size:
self._trigger_alert('POSITION_LIMIT_EXCEEDED', {
'current_ratio': position_ratio,
'limit': self.max_position_size
})
return False
return True
def _check_drawdown_limit(self, portfolio_state):
"""Check maximum drawdown limit"""
current_value = portfolio_state['total_value']
if current_value > self.max_portfolio_value:
self.max_portfolio_value = current_value
drawdown = (self.max_portfolio_value - current_value) / self.max_portfolio_value
if drawdown > self.max_drawdown:
self._trigger_alert('DRAWDOWN_LIMIT_EXCEEDED', {
'current_drawdown': drawdown,
'limit': self.max_drawdown
})
return False
return True
def _check_daily_loss_limit(self):
"""Check daily loss limit"""
if self.daily_pnl < -self.max_daily_loss:
self._trigger_alert('DAILY_LOSS_LIMIT_EXCEEDED', {
'daily_pnl': self.daily_pnl,
'limit': -self.max_daily_loss
})
return False
return True
def _check_var_limit(self, portfolio_state):
"""Check Value at Risk limit"""
# Calculate VaR (simplified)
var_95 = self._calculate_var(portfolio_state)
if var_95 > self.var_limit:
self._trigger_alert('VAR_LIMIT_EXCEEDED', {
'var_95': var_95,
'limit': self.var_limit
})
return False
return True
def _calculate_var(self, portfolio_state, confidence=0.95):
"""Calculate Value at Risk"""
# Simplified VaR calculation
# In production, this would use historical returns or Monte Carlo simulation
portfolio_volatility = portfolio_state.get('volatility', 0.02)
portfolio_value = portfolio_state['total_value']
# VaR = z_score * volatility * portfolio_value
z_score = 1.645 # For 95% confidence
var = z_score * portfolio_volatility * portfolio_value
return var
def _trigger_alert(self, alert_type, data):
"""Trigger risk alert"""
alert = {
'timestamp': datetime.now(),
'type': alert_type,
'data': data,
'severity': 'HIGH'
}
self.alerts.append(alert)
# Log alert
logger.error(f"Risk alert triggered: {alert_type}", extra=data)
# Notify risk team
self._notify_risk_team(alert)
def _notify_risk_team(self, alert):
"""Notify risk management team"""
# Send email/SMS notification
notification_service.send_alert(alert)
# Update dashboard
dashboard.update_risk_status(alert)
def handle_error(self, error):
"""Handle system errors"""
# Log error
logger.error(f"System error in risk manager: {error}")
# Trigger emergency procedures
self._emergency_procedures(error)
def _emergency_procedures(self, error):
"""Emergency risk procedures"""
# Close all positions
emergency_close = {
'action': 'emergency_close',
'reason': str(error),
'timestamp': datetime.now()
}
# Notify execution engine
execution_engine.emergency_close()
# Notify management
self._notify_management(emergency_close)
Circuit Breakers
Circuit Breaker Implementation:
class CircuitBreaker:
def __init__(self, failure_threshold=5, timeout=300):
self.failure_threshold = failure_threshold
self.timeout = timeout
self.failure_count = 0
self.last_failure_time = None
self.state = 'CLOSED' # CLOSED, OPEN, HALF_OPEN
def call(self, func, *args, **kwargs):
"""Execute function with circuit breaker protection"""
if self.state == 'OPEN':
if time.time() - self.last_failure_time > self.timeout:
self.state = 'HALF_OPEN'
else:
raise Exception("Circuit breaker is OPEN")
try:
result = func(*args, **kwargs)
self._on_success()
return result
except Exception as e:
self._on_failure()
raise e
def _on_success(self):
"""Handle successful call"""
self.failure_count = 0
self.state = 'CLOSED'
def _on_failure(self):
"""Handle failed call"""
self.failure_count += 1
self.last_failure_time = time.time()
if self.failure_count >= self.failure_threshold:
self.state = 'OPEN'
logger.error(f"Circuit breaker opened after {self.failure_count} failures")
4. Monitoring and Alerting
System Monitoring
class ProductionMonitoring:
def __init__(self, config):
self.config = config
# Metrics storage
self.metrics = {}
self.alerts = []
# External services
self.prometheus = PrometheusClient()
self.grafana = GrafanaClient()
self.slack = SlackClient(config['slack_webhook'])
def start_monitoring(self):
"""Start monitoring system"""
# Start metric collection
self._start_metric_collection()
# Start alert processing
self._start_alert_processing()
# Start health checks
self._start_health_checks()
def _start_metric_collection(self):
"""Start collecting system metrics"""
# Trading metrics
self._collect_trading_metrics()
# System metrics
self._collect_system_metrics()
# Model metrics
self._collect_model_metrics()
def _collect_trading_metrics(self):
"""Collect trading performance metrics"""
metrics = {
'portfolio_value': self.portfolio_state['total_value'],
'daily_pnl': self.daily_pnl,
'positions': len(self.portfolio_state['positions']),
'trades_today': self.trades_today,
'win_rate': self.calculate_win_rate()
}
# Send to Prometheus
for metric_name, value in metrics.items():
self.prometheus.gauge(f'trading_{metric_name}').set(value)
def _collect_system_metrics(self):
"""Collect system performance metrics"""
import psutil
metrics = {
'cpu_usage': psutil.cpu_percent(),
'memory_usage': psutil.virtual_memory().percent,
'disk_usage': psutil.disk_usage('/').percent,
'network_io': psutil.net_io_counters()._asdict()
}
# Send to Prometheus
for metric_name, value in metrics.items():
self.prometheus.gauge(f'system_{metric_name}').set(value)
def _collect_model_metrics(self):
"""Collect model performance metrics"""
metrics = {
'inference_time': self.model_server.get_avg_inference_time(),
'cache_hit_rate': self.model_server.get_cache_hit_rate(),
'prediction_confidence': self.model_server.get_avg_confidence(),
'model_accuracy': self.calculate_model_accuracy()
}
# Send to Prometheus
for metric_name, value in metrics.items():
self.prometheus.gauge(f'model_{metric_name}').set(value)
def log_trade(self, trade_result):
"""Log trade execution"""
trade_log = {
'timestamp': datetime.now(),
'symbol': trade_result['symbol'],
'action': trade_result['action'],
'quantity': trade_result['quantity'],
'price': trade_result['price'],
'execution_time': trade_result['execution_time'],
'slippage': trade_result['slippage']
}
# Store in database
self.database.store_trade(trade_log)
# Send to monitoring
self.prometheus.counter('trades_executed').inc()
# Check for alerts
self._check_trade_alerts(trade_log)
def _check_trade_alerts(self, trade_log):
"""Check for trade-related alerts"""
# High slippage alert
if trade_log['slippage'] > 0.01: # 1% slippage
self._send_alert('HIGH_SLIPPAGE', {
'symbol': trade_log['symbol'],
'slippage': trade_log['slippage']
})
# Slow execution alert
if trade_log['execution_time'] > 5.0: # 5 seconds
self._send_alert('SLOW_EXECUTION', {
'symbol': trade_log['symbol'],
'execution_time': trade_log['execution_time']
})
def _send_alert(self, alert_type, data):
"""Send alert to monitoring systems"""
alert = {
'timestamp': datetime.now(),
'type': alert_type,
'data': data,
'severity': self._get_alert_severity(alert_type)
}
# Store alert
self.alerts.append(alert)
# Send to Slack
self.slack.send_alert(alert)
# Send to Grafana
self.grafana.create_annotation(alert)
# Log alert
logger.warning(f"Alert triggered: {alert_type}", extra=data)
def _get_alert_severity(self, alert_type):
"""Get alert severity level"""
severity_map = {
'HIGH_SLIPPAGE': 'WARNING',
'SLOW_EXECUTION': 'WARNING',
'POSITION_LIMIT_EXCEEDED': 'CRITICAL',
'DRAWDOWN_LIMIT_EXCEEDED': 'CRITICAL',
'VAR_LIMIT_EXCEEDED': 'CRITICAL'
}
return severity_map.get(alert_type, 'INFO')
Performance Dashboard
Real-time Dashboard:
class TradingDashboard:
def __init__(self, config):
self.config = config
self.metrics = {}
self.alerts = []
def update_dashboard(self, portfolio_state, market_state, system_state):
"""Update dashboard with latest data"""
# Portfolio metrics
self.metrics['portfolio'] = {
'total_value': portfolio_state['total_value'],
'daily_pnl': portfolio_state['daily_pnl'],
'positions': portfolio_state['positions'],
'cash': portfolio_state['cash']
}
# Market metrics
self.metrics['market'] = {
'current_prices': market_state['prices'],
'volatility': market_state['volatility'],
'volume': market_state['volume']
}
# System metrics
self.metrics['system'] = {
'cpu_usage': system_state['cpu_usage'],
'memory_usage': system_state['memory_usage'],
'model_inference_time': system_state['inference_time']
}
# Update web dashboard
self._update_web_dashboard()
def _update_web_dashboard(self):
"""Update web-based dashboard"""
# Send metrics to web dashboard
dashboard_api.update_metrics(self.metrics)
# Update charts
dashboard_api.update_charts(self.metrics)
# Update alerts
dashboard_api.update_alerts(self.alerts)
5. Deployment Strategies
Blue-Green Deployment
class BlueGreenDeployment:
def __init__(self, config):
self.config = config
self.blue_version = None
self.green_version = None
self.current_version = 'blue'
def deploy_new_version(self, new_model_path):
"""Deploy new model version using blue-green strategy"""
# Determine target environment
target_version = 'green' if self.current_version == 'blue' else 'blue'
# Deploy new model to target environment
self._deploy_to_environment(target_version, new_model_path)
# Run validation tests
if self._validate_deployment(target_version):
# Switch traffic to new version
self._switch_traffic(target_version)
# Keep old version for rollback
self._keep_old_version_for_rollback()
return True
else:
# Rollback deployment
self._rollback_deployment(target_version)
return False
def _deploy_to_environment(self, version, model_path):
"""Deploy model to specific environment"""
# Load new model
new_model = self._load_model(model_path)
# Start new model server
if version == 'blue':
self.blue_version = ModelServer(self.config, new_model)
self.blue_version.start()
else:
self.green_version = ModelServer(self.config, new_model)
self.green_version.start()
def _validate_deployment(self, version):
"""Validate new deployment"""
# Run smoke tests
smoke_tests = SmokeTests()
if not smoke_tests.run_tests(version):
return False
# Run performance tests
performance_tests = PerformanceTests()
if not performance_tests.run_tests(version):
return False
# Run trading simulation
trading_simulation = TradingSimulation()
if not trading_simulation.run_simulation(version):
return False
return True
def _switch_traffic(self, new_version):
"""Switch traffic to new version"""
# Update load balancer
load_balancer.switch_traffic(new_version)
# Update current version
self.current_version = new_version
logger.info(f"Traffic switched to {new_version} version")
def _rollback_deployment(self, failed_version):
"""Rollback failed deployment"""
# Stop failed version
if failed_version == 'blue':
self.blue_version.stop()
self.blue_version = None
else:
self.green_version.stop()
self.green_version = None
logger.warning(f"Rolled back {failed_version} version deployment")
Canary Deployment
Gradual Rollout Strategy:
class CanaryDeployment:
def __init__(self, config):
self.config = config
self.canary_traffic_percentage = 0
self.canary_version = None
self.stable_version = None
def deploy_canary(self, new_model_path):
"""Deploy new version as canary"""
# Deploy canary version
self.canary_version = ModelServer(self.config, new_model_path)
self.canary_version.start()
# Start with 5% traffic
self.canary_traffic_percentage = 5
load_balancer.set_canary_traffic(5)
# Start monitoring
self._start_canary_monitoring()
def _start_canary_monitoring(self):
"""Start monitoring canary deployment"""
# Monitor key metrics
metrics_to_monitor = [
'error_rate',
'latency',
'trading_performance',
'model_accuracy'
]
for metric in metrics_to_monitor:
self._monitor_metric(metric)
def _monitor_metric(self, metric_name):
"""Monitor specific metric for canary"""
# Get canary and stable metrics
canary_metric = self._get_canary_metric(metric_name)
stable_metric = self._get_stable_metric(metric_name)
# Compare metrics
if self._is_metric_degraded(canary_metric, stable_metric):
# Rollback canary
self._rollback_canary()
return
# If metrics are good, increase traffic
if self._should_increase_traffic():
self._increase_canary_traffic()
def _should_increase_traffic(self):
"""Determine if canary traffic should be increased"""
# Check if canary has been stable for required time
if self.canary_stable_time < self.config['min_stable_time']:
return False
# Check if metrics are within acceptable range
return self._are_metrics_acceptable()
def _increase_canary_traffic(self):
"""Increase canary traffic percentage"""
# Increase by 10% each time
new_percentage = min(self.canary_traffic_percentage + 10, 100)
if new_percentage == 100:
# Full rollout complete
self._complete_rollout()
else:
# Increase traffic
load_balancer.set_canary_traffic(new_percentage)
self.canary_traffic_percentage = new_percentage
self.canary_stable_time = 0 # Reset timer
def _complete_rollout(self):
"""Complete canary rollout"""
# Make canary the stable version
self.stable_version = self.canary_version
self.canary_version = None
self.canary_traffic_percentage = 0
# Update load balancer
load_balancer.set_canary_traffic(0)
logger.info("Canary deployment completed successfully")