Team Ai
Modelpublic

ParallelLLC/algorithmic_trading

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
27likes32downloads
jupyter_widgets.py556 linesDownload Raw Back to ui
1"""2Jupyter Widgets UI for Algorithmic Trading System3 4Interactive notebook interface for:5- Data exploration and visualization6- Strategy development and testing7- Model training and evaluation8- Real-time trading simulation9"""10 11import ipywidgets as widgets12from IPython.display import display, HTML, clear_output13import plotly.graph_objects as go14import plotly.express as px15import pandas as pd16import numpy as np17import yaml18import os19import sys20from datetime import datetime, timedelta21from typing import Dict, Any, Optional22import asyncio23import threading24import time25 26# Add project root to path27sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))28 29from agentic_ai_system.main import load_config30from agentic_ai_system.data_ingestion import load_data, validate_data, add_technical_indicators31from agentic_ai_system.finrl_agent import FinRLAgent, FinRLConfig32from agentic_ai_system.alpaca_broker import AlpacaBroker33from agentic_ai_system.orchestrator import run_backtest, run_live_trading34 35class TradingJupyterUI:36    def __init__(self):37        self.config = None38        self.data = None39        self.alpaca_broker = None40        self.finrl_agent = None41        self.trading_active = False42        43        self.setup_widgets()44    45    def setup_widgets(self):46        """Setup all interactive widgets"""47        48        # Configuration widgets49        self.config_file = widgets.Text(50            value='config.yaml',51            description='Config File:',52            style={'description_width': '120px'}53        )54        55        self.load_config_btn = widgets.Button(56            description='Load Configuration',57            button_style='primary',58            icon='cog'59        )60        61        self.config_output = widgets.Output()62        63        # Data widgets64        self.data_source = widgets.Dropdown(65            options=['yahoo', 'csv', 'alpaca', 'synthetic'],66            value='yahoo',67            description='Data Source:',68            style={'description_width': '120px'}69        )70        71        self.symbol_input = widgets.Text(72            value='AAPL',73            description='Symbol:',74            style={'description_width': '120px'}75        )76        77        self.timeframe_input = widgets.Dropdown(78            options=['1m', '5m', '15m', '1h', '1d'],79            value='1d',80            description='Timeframe:',81            style={'description_width': '120px'}82        )83        84        self.load_data_btn = widgets.Button(85            description='Load Data',86            button_style='success',87            icon='database'88        )89        90        self.data_output = widgets.Output()91        92        # Alpaca widgets93        self.alpaca_api_key = widgets.Password(94            description='API Key:',95            style={'description_width': '120px'}96        )97        98        self.alpaca_secret_key = widgets.Password(99            description='Secret Key:',100            style={'description_width': '120px'}101        )102        103        self.connect_alpaca_btn = widgets.Button(104            description='Connect to Alpaca',105            button_style='info',106            icon='link'107        )108        109        self.alpaca_output = widgets.Output()110        111        # FinRL widgets112        self.finrl_algorithm = widgets.Dropdown(113            options=['PPO', 'A2C', 'DDPG', 'TD3'],114            value='PPO',115            description='Algorithm:',116            style={'description_width': '120px'}117        )118        119        self.learning_rate = widgets.FloatSlider(120            value=0.0003,121            min=0.0001,122            max=0.01,123            step=0.0001,124            description='Learning Rate:',125            style={'description_width': '120px'},126            readout_format='.4f'127        )128        129        self.training_steps = widgets.IntSlider(130            value=100000,131            min=1000,132            max=1000000,133            step=1000,134            description='Training Steps:',135            style={'description_width': '120px'}136        )137        138        self.batch_size = widgets.Dropdown(139            options=[32, 64, 128, 256],140            value=64,141            description='Batch Size:',142            style={'description_width': '120px'}143        )144        145        self.start_training_btn = widgets.Button(146            description='Start Training',147            button_style='warning',148            icon='play'149        )150        151        self.finrl_output = widgets.Output()152        153        # Trading widgets154        self.capital_input = widgets.IntText(155            value=100000,156            description='Capital ($):',157            style={'description_width': '120px'}158        )159        160        self.order_size_input = widgets.IntText(161            value=10,162            description='Order Size:',163            style={'description_width': '120px'}164        )165        166        self.start_trading_btn = widgets.Button(167            description='Start Trading',168            button_style='danger',169            icon='rocket'170        )171        172        self.stop_trading_btn = widgets.Button(173            description='Stop Trading',174            button_style='danger',175            icon='stop'176        )177        178        self.trading_output = widgets.Output()179        180        # Backtesting widgets181        self.run_backtest_btn = widgets.Button(182            description='Run Backtest',183            button_style='primary',184            icon='chart-line'185        )186        187        self.backtest_output = widgets.Output()188        189        # Chart widgets190        self.chart_type = widgets.Dropdown(191            options=['Candlestick', 'Line', 'Volume', 'Technical Indicators'],192            value='Candlestick',193            description='Chart Type:',194            style={'description_width': '120px'}195        )196        197        self.chart_output = widgets.Output()198        199        # Setup callbacks200        self.load_config_btn.on_click(self.on_load_config)201        self.load_data_btn.on_click(self.on_load_data)202        self.connect_alpaca_btn.on_click(self.on_connect_alpaca)203        self.start_training_btn.on_click(self.on_start_training)204        self.start_trading_btn.on_click(self.on_start_trading)205        self.stop_trading_btn.on_click(self.on_stop_trading)206        self.run_backtest_btn.on_click(self.on_run_backtest)207        self.chart_type.observe(self.on_chart_type_change, names='value')208    209    def on_load_config(self, b):210        """Handle configuration loading"""211        with self.config_output:212            clear_output()213            try:214                self.config = load_config(self.config_file.value)215                print(f"✅ Configuration loaded from {self.config_file.value}")216                print(f"Symbol: {self.config['trading']['symbol']}")217                print(f"Capital: ${self.config['trading']['capital']:,}")218                print(f"Timeframe: {self.config['trading']['timeframe']}")219                print(f"Broker: {self.config['execution']['broker_api']}")220            except Exception as e:221                print(f"❌ Error loading configuration: {e}")222    223    def on_load_data(self, b):224        """Handle data loading"""225        with self.data_output:226            clear_output()227            try:228                if self.config:229                    # Update config with widget values230                    self.config['data_source']['type'] = self.data_source.value231                    self.config['trading']['symbol'] = self.symbol_input.value232                    self.config['trading']['timeframe'] = self.timeframe_input.value233                    234                    print(f"Loading data for {self.symbol_input.value}...")235                    self.data = load_data(self.config)236                    237                    if self.data is not None and not self.data.empty:238                        print(f"✅ Loaded {len(self.data)} data points")239                        print(f"Date range: {self.data['timestamp'].min()} to {self.data['timestamp'].max()}")240                        print(f"Price range: ${self.data['close'].min():.2f} - ${self.data['close'].max():.2f}")241                        242                        # Add technical indicators243                        self.data = add_technical_indicators(self.data)244                        print(f"✅ Added technical indicators")245                        246                        # Update chart247                        self.update_chart()248                    else:249                        print("❌ Failed to load data")250                else:251                    print("⚠️ Please load configuration first")252            except Exception as e:253                print(f"❌ Error loading data: {e}")254    255    def on_connect_alpaca(self, b):256        """Handle Alpaca connection"""257        with self.alpaca_output:258            clear_output()259            try:260                if self.alpaca_api_key.value and self.alpaca_secret_key.value:261                    # Update config with API keys262                    if self.config:263                        self.config['alpaca']['api_key'] = self.alpaca_api_key.value264                        self.config['alpaca']['secret_key'] = self.alpaca_secret_key.value265                        self.config['execution']['broker_api'] = 'alpaca_paper'266                        267                        print("Connecting to Alpaca...")268                        self.alpaca_broker = AlpacaBroker(self.config)269                        270                        account_info = self.alpaca_broker.get_account_info()271                        if account_info:272                            print("✅ Connected to Alpaca")273                            print(f"Account ID: {account_info['account_id']}")274                            print(f"Status: {account_info['status']}")275                            print(f"Buying Power: ${account_info['buying_power']:,.2f}")276                            print(f"Portfolio Value: ${account_info['portfolio_value']:,.2f}")277                        else:278                            print("❌ Failed to connect to Alpaca")279                    else:280                        print("⚠️ Please load configuration first")281                else:282                    print("⚠️ Please enter Alpaca API credentials")283            except Exception as e:284                print(f"❌ Error connecting to Alpaca: {e}")285    286    def on_start_training(self, b):287        """Handle FinRL training"""288        with self.finrl_output:289            clear_output()290            try:291                if self.data is not None:292                    print("Starting FinRL training...")293                    294                    # Create FinRL config295                    finrl_config = FinRLConfig(296                        algorithm=self.finrl_algorithm.value,297                        learning_rate=self.learning_rate.value,298                        batch_size=self.batch_size.value,299                        buffer_size=1000000,300                        learning_starts=100,301                        gamma=0.99,302                        tau=0.005,303                        train_freq=1,304                        gradient_steps=1,305                        verbose=1,306                        tensorboard_log='logs/finrl_tensorboard'307                    )308                    309                    # Initialize agent310                    self.finrl_agent = FinRLAgent(finrl_config)311                    312                    # Train the agent313                    result = self.finrl_agent.train(314                        data=self.data,315                        config=self.config,316                        total_timesteps=self.training_steps.value,317                        use_real_broker=False318                    )319                    320                    if result['success']:321                        print("✅ Training completed successfully!")322                        print(f"Algorithm: {result['algorithm']}")323                        print(f"Timesteps: {result['total_timesteps']:,}")324                        print(f"Model saved: {result['model_path']}")325                    else:326                        print("❌ Training failed")327                else:328                    print("⚠️ Please load data first")329            except Exception as e:330                print(f"❌ Error during training: {e}")331    332    def on_start_trading(self, b):333        """Handle trading start"""334        with self.trading_output:335            clear_output()336            try:337                if self.config and self.alpaca_broker:338                    print("Starting live trading...")339                    self.trading_active = True340                    341                    # Update config with widget values342                    self.config['trading']['capital'] = self.capital_input.value343                    self.config['execution']['order_size'] = self.order_size_input.value344                    345                    # Start trading in background thread346                    def run_trading():347                        try:348                            run_live_trading(self.config, self.data)349                        except Exception as e:350                            print(f"Trading error: {e}")351                    352                    trading_thread = threading.Thread(target=run_trading)353                    trading_thread.daemon = True354                    trading_thread.start()355                    356                    print("✅ Live trading started")357                else:358                    print("⚠️ Please load configuration and connect to Alpaca first")359            except Exception as e:360                print(f"❌ Error starting trading: {e}")361    362    def on_stop_trading(self, b):363        """Handle trading stop"""364        with self.trading_output:365            clear_output()366            self.trading_active = False367            print("✅ Trading stopped")368    369    def on_run_backtest(self, b):370        """Handle backtesting"""371        with self.backtest_output:372            clear_output()373            try:374                if self.config and self.data is not None:375                    print("Running backtest...")376                    377                    # Update config with widget values378                    self.config['trading']['capital'] = self.capital_input.value379                    380                    result = run_backtest(self.config, self.data)381                    382                    if result['success']:383                        print("✅ Backtest completed")384                        print(f"Total Return: {result['total_return']:.2%}")385                        print(f"Sharpe Ratio: {result['sharpe_ratio']:.2f}")386                        print(f"Max Drawdown: {result['max_drawdown']:.2%}")387                        print(f"Total Trades: {result['total_trades']}")388                    else:389                        print("❌ Backtest failed")390                else:391                    print("⚠️ Please load configuration and data first")392            except Exception as e:393                print(f"❌ Error during backtest: {e}")394    395    def on_chart_type_change(self, change):396        """Handle chart type change"""397        if self.data is not None:398            self.update_chart()399    400    def update_chart(self):401        """Update the chart display"""402        with self.chart_output:403            clear_output()404            405            if self.data is None:406                return407            408            if self.chart_type.value == "Candlestick":409                fig = go.Figure(data=[go.Candlestick(410                    x=self.data['timestamp'],411                    open=self.data['open'],412                    high=self.data['high'],413                    low=self.data['low'],414                    close=self.data['close']415                )])416                fig.update_layout(417                    title=f"{self.config['trading']['symbol']} Candlestick Chart",418                    xaxis_title="Date",419                    yaxis_title="Price ($)",420                    height=500421                )422                display(fig)423            424            elif self.chart_type.value == "Line":425                fig = px.line(self.data, x='timestamp', y='close',426                             title=f"{self.config['trading']['symbol']} Price Chart")427                fig.update_layout(height=500)428                display(fig)429            430            elif self.chart_type.value == "Volume":431                fig = go.Figure()432                fig.add_trace(go.Bar(433                    x=self.data['timestamp'],434                    y=self.data['volume'],435                    name='Volume'436                ))437                fig.update_layout(438                    title=f"{self.config['trading']['symbol']} Volume Chart",439                    xaxis_title="Date",440                    yaxis_title="Volume",441                    height=500442                )443                display(fig)444            445            elif self.chart_type.value == "Technical Indicators":446                fig = go.Figure()447                448                # Add price449                fig.add_trace(go.Scatter(450                    x=self.data['timestamp'],451                    y=self.data['close'],452                    name='Close Price',453                    line=dict(color='blue')454                ))455                456                # Add moving averages if available457                if 'sma_20' in self.data.columns:458                    fig.add_trace(go.Scatter(459                        x=self.data['timestamp'],460                        y=self.data['sma_20'],461                        name='SMA 20',462                        line=dict(color='orange')463                    ))464                465                if 'sma_50' in self.data.columns:466                    fig.add_trace(go.Scatter(467                        x=self.data['timestamp'],468                        y=self.data['sma_50'],469                        name='SMA 50',470                        line=dict(color='red')471                    ))472                473                fig.update_layout(474                    title=f"{self.config['trading']['symbol']} Technical Indicators",475                    xaxis_title="Date",476                    yaxis_title="Price ($)",477                    height=500478                )479                display(fig)480    481    def display_interface(self):482        """Display the complete Jupyter interface"""483        484        # Header485        display(HTML("""486        <div style="text-align: center; margin-bottom: 20px;">487            <h1>🤖 Algorithmic Trading System</h1>488            <p>Interactive Jupyter Interface for Trading Analysis</p>489        </div>490        """))491        492        # Configuration section493        display(HTML("<h2>⚙️ Configuration</h2>"))494        config_widgets = widgets.VBox([495            widgets.HBox([self.config_file, self.load_config_btn]),496            self.config_output497        ])498        display(config_widgets)499        500        # Data section501        display(HTML("<h2>📥 Data Management</h2>"))502        data_widgets = widgets.VBox([503            widgets.HBox([self.data_source, self.symbol_input, self.timeframe_input]),504            widgets.HBox([self.load_data_btn]),505            self.data_output506        ])507        display(data_widgets)508        509        # Alpaca section510        display(HTML("<h2>🏦 Alpaca Integration</h2>"))511        alpaca_widgets = widgets.VBox([512            widgets.HBox([self.alpaca_api_key, self.alpaca_secret_key]),513            widgets.HBox([self.connect_alpaca_btn]),514            self.alpaca_output515        ])516        display(alpaca_widgets)517        518        # FinRL section519        display(HTML("<h2>🧠 FinRL Training</h2>"))520        finrl_widgets = widgets.VBox([521            widgets.HBox([self.finrl_algorithm, self.learning_rate]),522            widgets.HBox([self.training_steps, self.batch_size]),523            widgets.HBox([self.start_training_btn]),524            self.finrl_output525        ])526        display(finrl_widgets)527        528        # Trading section529        display(HTML("<h2>🎯 Trading Controls</h2>"))530        trading_widgets = widgets.VBox([531            widgets.HBox([self.capital_input, self.order_size_input]),532            widgets.HBox([self.start_trading_btn, self.stop_trading_btn]),533            self.trading_output534        ])535        display(trading_widgets)536        537        # Backtesting section538        display(HTML("<h2>📊 Backtesting</h2>"))539        backtest_widgets = widgets.VBox([540            widgets.HBox([self.run_backtest_btn]),541            self.backtest_output542        ])543        display(backtest_widgets)544        545        # Chart section546        display(HTML("<h2>📈 Data Visualization</h2>"))547        chart_widgets = widgets.VBox([548            widgets.HBox([self.chart_type]),549            self.chart_output550        ])551        display(chart_widgets)552 553def create_jupyter_interface():554    """Create and return the Jupyter interface"""555    ui = TradingJupyterUI()556    return ui