ParallelLLC/algorithmic_trading
2732
1#!/usr/bin/env python32"""3FinRL Demo Script4 5This script demonstrates the integration of FinRL with the algorithmic trading system.6It shows how to train a reinforcement learning agent and use it for trading decisions.7"""8 9import os10import sys11import yaml12import pandas as pd13import numpy as np14import matplotlib.pyplot as plt15import seaborn as sns16from datetime import datetime, timedelta17import logging18 19# Add the project root to the path20sys.path.append(os.path.dirname(os.path.abspath(__file__)))21 22from agentic_ai_system.finrl_agent import FinRLAgent, FinRLConfig, create_finrl_agent_from_config23from agentic_ai_system.synthetic_data_generator import SyntheticDataGenerator24from agentic_ai_system.logger_config import setup_logging25 26 27def load_config(config_path: str = 'config.yaml') -> dict:28 """Load configuration from YAML file"""29 with open(config_path, 'r') as file:30 return yaml.safe_load(file)31 32# Setup logging33config = load_config()34setup_logging(config)35logger = logging.getLogger(__name__)36 37 38def generate_training_data(config: dict) -> pd.DataFrame:39 """Generate synthetic data for training"""40 logger.info("Generating synthetic training data")41 42 generator = SyntheticDataGenerator(config)43 44 # Generate training data (longer period)45 train_data = generator.generate_ohlcv_data(46 symbol='AAPL',47 start_date='2023-01-01',48 end_date='2023-12-31',49 frequency='1H'50 )51 52 # Add technical indicators53 train_data['sma_20'] = train_data['close'].rolling(window=20).mean()54 train_data['sma_50'] = train_data['close'].rolling(window=50).mean()55 train_data['rsi'] = calculate_rsi(train_data['close'])56 bb_upper, bb_lower = calculate_bollinger_bands(train_data['close'])57 train_data['bb_upper'] = bb_upper58 train_data['bb_lower'] = bb_lower59 train_data['macd'] = calculate_macd(train_data['close'])60 61 # Fill NaN values62 train_data = train_data.fillna(method='bfill').fillna(0)63 64 logger.info(f"Generated {len(train_data)} training samples")65 return train_data66 67 68def generate_test_data(config: dict) -> pd.DataFrame:69 """Generate synthetic data for testing"""70 logger.info("Generating synthetic test data")71 72 generator = SyntheticDataGenerator(config)73 74 # Generate test data (shorter period)75 test_data = generator.generate_ohlcv_data(76 symbol='AAPL',77 start_date='2024-01-01',78 end_date='2024-03-31',79 frequency='1H'80 )81 82 # Add technical indicators83 test_data['sma_20'] = test_data['close'].rolling(window=20).mean()84 test_data['sma_50'] = test_data['close'].rolling(window=50).mean()85 test_data['rsi'] = calculate_rsi(test_data['close'])86 bb_upper, bb_lower = calculate_bollinger_bands(test_data['close'])87 test_data['bb_upper'] = bb_upper88 test_data['bb_lower'] = bb_lower89 test_data['macd'] = calculate_macd(test_data['close'])90 91 # Fill NaN values92 test_data = test_data.fillna(method='bfill').fillna(0)93 94 logger.info(f"Generated {len(test_data)} test samples")95 return test_data96 97 98def calculate_rsi(prices: pd.Series, period: int = 14) -> pd.Series:99 """Calculate RSI indicator"""100 delta = prices.diff()101 gain = (delta.where(delta > 0, 0)).rolling(window=period).mean()102 loss = (-delta.where(delta < 0, 0)).rolling(window=period).mean()103 rs = gain / loss104 rsi = 100 - (100 / (1 + rs))105 return rsi106 107 108def calculate_bollinger_bands(prices: pd.Series, period: int = 20, std_dev: int = 2):109 """Calculate Bollinger Bands"""110 sma = prices.rolling(window=period).mean()111 std = prices.rolling(window=period).std()112 upper_band = sma + (std * std_dev)113 lower_band = sma - (std * std_dev)114 return upper_band, lower_band115 116 117def calculate_macd(prices: pd.Series, fast: int = 12, slow: int = 26, signal: int = 9) -> pd.Series:118 """Calculate MACD indicator"""119 ema_fast = prices.ewm(span=fast).mean()120 ema_slow = prices.ewm(span=slow).mean()121 macd_line = ema_fast - ema_slow122 return macd_line123 124 125def train_finrl_agent(config: dict, train_data: pd.DataFrame, test_data: pd.DataFrame) -> FinRLAgent:126 """Train the FinRL agent"""127 logger.info("Starting FinRL agent training")128 129 # Create FinRL agent130 finrl_config = FinRLConfig(**{k: v for k, v in config['finrl'].items() if k in FinRLConfig.__dataclass_fields__})131 agent = FinRLAgent(finrl_config)132 133 # Train the agent134 training_result = agent.train(135 data=train_data,136 config=config,137 total_timesteps=config['finrl']['training']['total_timesteps']138 )139 140 logger.info(f"Training completed: {training_result}")141 142 # Save the model143 if config['finrl']['training']['save_best_model']:144 model_path = config['finrl']['training']['model_save_path']145 os.makedirs(os.path.dirname(model_path), exist_ok=True)146 agent.save_model(model_path)147 148 return agent149 150 151def evaluate_agent(agent: FinRLAgent, test_data: pd.DataFrame, config: dict) -> dict:152 """Evaluate the trained agent"""153 logger.info("Evaluating FinRL agent")154 155 # Evaluate on test data156 evaluation_results = agent.evaluate(test_data, config)157 158 logger.info(f"Evaluation results: {evaluation_results}")159 160 return evaluation_results161 162 163def generate_predictions(agent: FinRLAgent, test_data: pd.DataFrame, config: dict) -> list:164 """Generate trading predictions"""165 logger.info("Generating trading predictions")166 167 prediction_results = agent.predict(test_data, config)168 169 if prediction_results['success']:170 predictions = prediction_results['actions']171 logger.info(f"Generated {len(predictions)} predictions")172 return predictions173 else:174 logger.error(f"Prediction failed: {prediction_results['error']}")175 return []176 177 178def plot_results(test_data: pd.DataFrame, predictions: list, evaluation_results: dict):179 """Plot trading results"""180 logger.info("Creating visualization plots")181 182 # Create figure with subplots183 fig, axes = plt.subplots(3, 1, figsize=(15, 12))184 185 # Plot 1: Price and predictions186 axes[0].plot(test_data.index, test_data['close'], label='Close Price', alpha=0.7)187 188 # Mark buy/sell signals only if predictions are available189 if predictions:190 buy_signals = [i for i, pred in enumerate(predictions) if pred == 2]191 sell_signals = [i for i, pred in enumerate(predictions) if pred == 0]192 193 if buy_signals:194 axes[0].scatter(test_data.index[buy_signals], test_data['close'].iloc[buy_signals], 195 color='green', marker='^', s=100, label='Buy Signal', alpha=0.8)196 if sell_signals:197 axes[0].scatter(test_data.index[sell_signals], test_data['close'].iloc[sell_signals], 198 color='red', marker='v', s=100, label='Sell Signal', alpha=0.8)199 200 axes[0].set_title('Price Action and Trading Signals')201 axes[0].set_ylabel('Price')202 axes[0].legend()203 axes[0].grid(True, alpha=0.3)204 205 # Plot 2: Technical indicators206 axes[1].plot(test_data.index, test_data['close'], label='Close Price', alpha=0.7)207 axes[1].plot(test_data.index, test_data['sma_20'], label='SMA 20', alpha=0.7)208 axes[1].plot(test_data.index, test_data['sma_50'], label='SMA 50', alpha=0.7)209 axes[1].plot(test_data.index, test_data['bb_upper'], label='BB Upper', alpha=0.5)210 axes[1].plot(test_data.index, test_data['bb_lower'], label='BB Lower', alpha=0.5)211 212 axes[1].set_title('Technical Indicators')213 axes[1].set_ylabel('Price')214 axes[1].legend()215 axes[1].grid(True, alpha=0.3)216 217 # Plot 3: RSI218 axes[2].plot(test_data.index, test_data['rsi'], label='RSI', color='purple')219 axes[2].axhline(y=70, color='r', linestyle='--', alpha=0.5, label='Overbought')220 axes[2].axhline(y=30, color='g', linestyle='--', alpha=0.5, label='Oversold')221 axes[2].set_title('RSI Indicator')222 axes[2].set_ylabel('RSI')223 axes[2].set_xlabel('Time')224 axes[2].legend()225 axes[2].grid(True, alpha=0.3)226 227 plt.tight_layout()228 229 # Save plot230 os.makedirs('plots', exist_ok=True)231 plt.savefig('plots/finrl_trading_results.png', dpi=300, bbox_inches='tight')232 plt.show()233 234 logger.info("Plots saved to plots/finrl_trading_results.png")235 236 237def print_summary(evaluation_results: dict, predictions: list):238 """Print trading summary"""239 print("\n" + "="*60)240 print("FINRL TRADING SYSTEM SUMMARY")241 print("="*60)242 243 if evaluation_results.get('success', False):244 print(f"Algorithm: {evaluation_results.get('algorithm', 'Unknown')}")245 print(f"Total Return: {evaluation_results.get('total_return', 0):.2%}")246 print(f"Final Portfolio Value: ${evaluation_results.get('final_portfolio_value', 0):,.2f}")247 print(f"Total Reward: {evaluation_results.get('total_reward', 0):.4f}")248 print(f"Sharpe Ratio: {evaluation_results.get('sharpe_ratio', 0):.4f}")249 print(f"Number of Trading Steps: {evaluation_results.get('steps', 0)}")250 print(f"Max Drawdown: {evaluation_results.get('max_drawdown', 0):.2%}")251 else:252 print(f"Evaluation failed: {evaluation_results.get('error', 'Unknown error')}")253 254 # Trading statistics255 if predictions:256 buy_signals = sum(1 for pred in predictions if pred == 2)257 sell_signals = sum(1 for pred in predictions if pred == 0)258 hold_signals = sum(1 for pred in predictions if pred == 1)259 260 print(f"\nTrading Signals:")261 print(f" Buy signals: {buy_signals}")262 print(f" Sell signals: {sell_signals}")263 print(f" Hold signals: {hold_signals}")264 print(f" Total signals: {len(predictions)}")265 else:266 print(f"\nNo trading predictions available")267 268 print("\n" + "="*60)269 270 271def main():272 """Main function to run the FinRL demo"""273 logger.info("Starting FinRL Demo")274 275 try:276 # Load configuration277 config = load_config()278 279 # Generate data280 train_data = generate_training_data(config)281 test_data = generate_test_data(config)282 283 # Train FinRL agent284 agent = train_finrl_agent(config, train_data, test_data)285 286 # Evaluate agent287 evaluation_results = evaluate_agent(agent, test_data, config)288 289 # Generate predictions290 predictions = generate_predictions(agent, test_data, config)291 292 # Create visualizations293 plot_results(test_data, predictions, evaluation_results)294 295 # Print summary296 print_summary(evaluation_results, predictions)297 298 logger.info("FinRL Demo completed successfully")299 300 except Exception as e:301 logger.error(f"Error in FinRL demo: {str(e)}")302 raise303 304 305if __name__ == "__main__":306 main() 