Team Ai
Modelpublic

ParallelLLC/algorithmic_trading

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
27likes32downloads
finrl_demo.py306 linesDownload Raw Back to root
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()