From 4bf7394a0a5a02bee40143eff88a58049e916f10 Mon Sep 17 00:00:00 2001 From: Florian Egger Date: Tue, 26 May 2026 13:51:02 +0200 Subject: [PATCH] initial commit --- .gitignore | 1 + benchmark.py | 176 ++++ config.py | 196 +++++ data/external/.gitkeep | 0 data/processed/.gitkeep | 0 data/raw/.gitkeep | 0 live_trading.py | 569 +++++++++++++ main.py | 156 ++++ models/.gitkeep | 0 notebooks/.gitkeep | 0 requirements.txt | 58 ++ src/__init__.py | 1 + src/amd/__init__.py | 1 + src/amd/optimizations.py | 349 ++++++++ src/data/__init__.py | 0 src/data/live_data.py | 908 ++++++++++++++++++++ src/data/news_processor.py | 36 + src/data/pipeline.py | 1131 +++++++++++++++++++++++++ src/data/sentiment.py | 31 + src/data/social_processor.py | 59 ++ src/evaluation/__init__.py | 0 src/evaluation/backtester.py | 61 ++ src/evaluation/intraday_backtester.py | 57 ++ src/evaluation/metrics.py | 128 +++ src/models/__init__.py | 0 src/models/gnn_model.py | 383 +++++++++ src/models/intraday_gnn.py | 220 +++++ src/models/trainer.py | 356 ++++++++ src/trading/__init__.py | 0 src/trading/broker.py | 38 + src/trading/ib_broker.py | 47 + src/trading/paper_broker.py | 82 ++ src/trading/real_time_trader.py | 91 ++ src/utils/__init__.py | 1 + src/utils/helpers.py | 46 + src/utils/memory_manager.py | 215 +++++ src/utils/text_processing.py | 12 + src/utils/visualization.py | 39 + src/web/__init__.py | 1 + src/web/api/__init__.py | 1 + src/web/api/dashboard.py | 84 ++ src/web/api/data_endpoints.py | 139 +++ src/web/api/models_endpoints.py | 193 +++++ src/web/api/trading_endpoints.py | 145 ++++ src/web/app.py | 193 +++++ src/web/services/__init__.py | 1 + src/web/services/state.py | 31 + src/web/static/css/style.css | 409 +++++++++ src/web/static/js/app.js | 419 +++++++++ src/web/templates/index.html | 268 ++++++ 50 files changed, 7332 insertions(+) create mode 100644 .gitignore create mode 100644 benchmark.py create mode 100644 config.py create mode 100644 data/external/.gitkeep create mode 100644 data/processed/.gitkeep create mode 100644 data/raw/.gitkeep create mode 100644 live_trading.py create mode 100644 main.py create mode 100644 models/.gitkeep create mode 100644 notebooks/.gitkeep create mode 100644 requirements.txt create mode 100644 src/__init__.py create mode 100644 src/amd/__init__.py create mode 100644 src/amd/optimizations.py create mode 100644 src/data/__init__.py create mode 100644 src/data/live_data.py create mode 100644 src/data/news_processor.py create mode 100644 src/data/pipeline.py create mode 100644 src/data/sentiment.py create mode 100644 src/data/social_processor.py create mode 100644 src/evaluation/__init__.py create mode 100644 src/evaluation/backtester.py create mode 100644 src/evaluation/intraday_backtester.py create mode 100644 src/evaluation/metrics.py create mode 100644 src/models/__init__.py create mode 100644 src/models/gnn_model.py create mode 100644 src/models/intraday_gnn.py create mode 100644 src/models/trainer.py create mode 100644 src/trading/__init__.py create mode 100644 src/trading/broker.py create mode 100644 src/trading/ib_broker.py create mode 100644 src/trading/paper_broker.py create mode 100644 src/trading/real_time_trader.py create mode 100644 src/utils/__init__.py create mode 100644 src/utils/helpers.py create mode 100644 src/utils/memory_manager.py create mode 100644 src/utils/text_processing.py create mode 100644 src/utils/visualization.py create mode 100644 src/web/__init__.py create mode 100644 src/web/api/__init__.py create mode 100644 src/web/api/dashboard.py create mode 100644 src/web/api/data_endpoints.py create mode 100644 src/web/api/models_endpoints.py create mode 100644 src/web/api/trading_endpoints.py create mode 100644 src/web/app.py create mode 100644 src/web/services/__init__.py create mode 100644 src/web/services/state.py create mode 100644 src/web/static/css/style.css create mode 100644 src/web/static/js/app.js create mode 100644 src/web/templates/index.html diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..5357087 --- /dev/null +++ b/.gitignore @@ -0,0 +1 @@ +implementation.md diff --git a/benchmark.py b/benchmark.py new file mode 100644 index 0000000..abc6aeb --- /dev/null +++ b/benchmark.py @@ -0,0 +1,176 @@ +import logging +import time + +import numpy as np +import torch +import torch.nn as nn + +from config import config +from src.amd.optimizations import AMDOptimizer +from src.models.intraday_gnn import IntradayGNN +from src.utils.memory_manager import MemoryManager + +# Configure logging +logging.basicConfig( + level=logging.INFO, + format="%(asctime)s - %(name)s - %(levelname)s - %(message)s", + handlers=[ + logging.FileHandler("benchmark_r9700.log"), + logging.StreamHandler(), + ], +) +logger = logging.getLogger(__name__) + + +def benchmark_model(): + """Benchmark the GNN model on AMD Radeon R9700 AI Pro""" + # Initialize memory manager + memory_manager = MemoryManager() + logger.info(memory_manager.get_memory_stats()) + + # Initialize AMD optimizer + amd_optimizer = AMDOptimizer() + + # Create a sample model + num_features = len(config.INTRADAY_FEATURES) + 5 # +5 for price features + model = IntradayGNN(num_features, config.SEQUENCE_LENGTH) + + # Optimize model for AMD GPU + model = amd_optimizer.optimize_model(model) + + # Create sample data + batch_size = config.BATCH_SIZE + sequence_length = config.SEQUENCE_LENGTH + num_stocks = 50 # Number of stocks in the graph + + # Create random data + x = torch.randn(num_stocks, sequence_length, num_features).to(config.DEVICE) + + # Create random edges + num_edges = 200 + edge_index = torch.randint(0, num_stocks, (2, num_edges)).to(config.DEVICE) + edge_attr = torch.randn(num_edges, 1).to(config.DEVICE) + + # Create target + y = torch.randn(num_stocks, 1).to(config.DEVICE) + + # Warm-up + logger.info("Warming up...") + for _ in range(10): + with torch.no_grad(): + _ = model((x, edge_index, edge_attr)) + + # Benchmark inference + logger.info("Benchmarking inference...") + start_time = time.time() + num_runs = 100 + + for _ in range(num_runs): + with torch.no_grad(): + _ = model((x, edge_index, edge_attr)) + + inference_time = (time.time() - start_time) / num_runs + logger.info(f"Average inference time: {inference_time:.6f} seconds") + + # Benchmark training + logger.info("Benchmarking training...") + model.train() + optimizer = torch.optim.Adam(model.parameters(), lr=config.LEARNING_RATE) + criterion = nn.MSELoss() + + start_time = time.time() + + for _ in range(num_runs): + optimizer.zero_grad() + out = model((x, edge_index, edge_attr)) + loss = criterion(out, y) + loss.backward() + optimizer.step() + + training_time = (time.time() - start_time) / num_runs + logger.info(f"Average training time: {training_time:.6f} seconds") + + # Memory usage + memory_info = memory_manager.check_memory() + logger.info( + f"GPU Memory Usage: {memory_info['allocated'] / 1024**3:.2f}GB / {memory_info['total'] / 1024**3:.2f}GB" + ) + + # Throughput + logger.info(f"Inference throughput: {1 / inference_time:.2f} samples/second") + logger.info(f"Training throughput: {1 / training_time:.2f} samples/second") + + # Detailed benchmark with different batch sizes + logger.info("\nDetailed benchmark with different configurations:") + + batch_sizes = [32, 64, 128, 256] + sequence_lengths = [30, 60, 120] + + for batch_size in batch_sizes: + for seq_len in sequence_lengths: + # Create data for this configuration + x = torch.randn(num_stocks, seq_len, num_features).to(config.DEVICE) + edge_index = torch.randint(0, num_stocks, (2, num_edges)).to(config.DEVICE) + edge_attr = torch.randn(num_edges, 1).to(config.DEVICE) + y = torch.randn(num_stocks, 1).to(config.DEVICE) + + # Benchmark inference + start_time = time.time() + for _ in range(10): # Fewer runs for detailed benchmark + with torch.no_grad(): + _ = model((x, edge_index, edge_attr)) + inf_time = (time.time() - start_time) / 10 + + # Benchmark training + start_time = time.time() + for _ in range(10): + optimizer.zero_grad() + out = model((x, edge_index, edge_attr)) + loss = criterion(out, y) + loss.backward() + optimizer.step() + train_time = (time.time() - start_time) / 10 + + logger.info( + f"Batch: {batch_size}, Seq Len: {seq_len}, " + f"Inf Time: {inf_time:.6f}s, Train Time: {train_time:.6f}s, " + f"Inf Tput: {1 / inf_time:.2f} samples/s, Train Tput: {1 / train_time:.2f} samples/s" + ) + + # Memory benchmark + logger.info("\nMemory benchmark:") + + # Test different model sizes + hidden_channels_list = [64, 128, 256, 512] + + for hidden_channels in hidden_channels_list: + # Create a model with this configuration + model = IntradayGNN(num_features, config.SEQUENCE_LENGTH) + model.feature_processor = nn.Sequential( + nn.Linear(num_features, hidden_channels), + nn.SiLU(), + nn.Linear(hidden_channels, hidden_channels), + nn.LayerNorm(hidden_channels), + ) + model.linear = nn.Linear(hidden_channels, 1) + + # Optimize model + model = amd_optimizer.optimize_model(model) + + # Estimate memory usage + estimated_memory = memory_manager.estimate_model_memory(model) + logger.info( + f"Hidden Channels: {hidden_channels}, Estimated Memory: {estimated_memory / 1024**3:.2f}GB" + ) + + # Clean up + del model + memory_manager.empty_cache() + + # Final memory stats + logger.info("\nFinal Memory Stats:") + logger.info(memory_manager.get_memory_stats()) + + +if __name__ == "__main__": + benchmark_model() diff --git a/config.py b/config.py new file mode 100644 index 0000000..722326d --- /dev/null +++ b/config.py @@ -0,0 +1,196 @@ +import os +from datetime import datetime, timedelta + +import torch + + +class Config: + # Project settings + PROJECT_NAME = "StockGNN_R9700" + VERSION = "1.0.0" + + # Data directories + BASE_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + DATA_DIR = os.path.join(BASE_DIR, "data") + RAW_DATA_DIR = os.path.join(DATA_DIR, "raw") + PROCESSED_DATA_DIR = os.path.join(DATA_DIR, "processed") + EXTERNAL_DATA_DIR = os.path.join(DATA_DIR, "external") + MODEL_DIR = os.path.join(BASE_DIR, "models") + + # Ensure directories exist + os.makedirs(RAW_DATA_DIR, exist_ok=True) + os.makedirs(PROCESSED_DATA_DIR, exist_ok=True) + os.makedirs(EXTERNAL_DATA_DIR, exist_ok=True) + os.makedirs(MODEL_DIR, exist_ok=True) + + # Stock universe settings + INITIAL_TICKERS = [ + "AAPL", + "MSFT", + "GOOGL", + "AMZN", + "META", + "TSLA", + "NVDA", + "JPM", + "V", + "WMT", + "PG", + "DIS", + "NFLX", + "ADBE", + "PYPL", + "INTC", + "CSCO", + "PEP", + "KO", + "XOM", + "BAC", + "VZ", + "T", + "CRM", + "CMCSA", + "PFE", + "NKE", + "MRK", + "CVX", + "HD", + ] + INDEX_TICKER = "^GSPC" # S&P 500 + DELISTED_TICKERS_FILE = os.path.join(EXTERNAL_DATA_DIR, "delisted_stocks.csv") + + # Date settings + START_DATE = "2015-01-01" + END_DATE = datetime.now().strftime("%Y-%m-%d") + TRAIN_END_DATE = "2022-12-31" + VAL_END_DATE = "2023-06-30" + TEST_END_DATE = END_DATE + + # AMD GPU settings (Radeon R9700 AI Pro - 32GB) + DEVICE = "cuda" if torch.cuda.is_available() else "cpu" + AMD_GPU = True + GPU_MEMORY_LIMIT = 0.9 # Use 90% of 32GB = 28.8GB + ROCM_OPT_LEVEL = "O2" # Optimization level for ROCm ('O0', 'O1', 'O2') + PIN_MEMORY = True # Enable pinned memory for faster data transfer + MIXED_PRECISION = True # Enable mixed precision training + PRECISION = "bf16" # 'fp16' or 'bf16' for mixed precision + + # Model settings + MODEL_NAME = "stock_gnn_r9700" + HIDDEN_CHANNELS = 128 # Increased for R9700's compute power + NUM_HEADS = 16 # Increased number of attention heads + DROPOUT = 0.3 # Reduced dropout for better GPU utilization + LEARNING_RATE = 0.0005 # Lower learning rate for stability + EPOCHS = 200 # More epochs with larger batch sizes + BATCH_SIZE = 128 # Larger batch size for R9700's memory + SEQUENCE_LENGTH = 60 # Longer sequences with more memory + PREDICTION_HORIZON = 10 # Number of steps to predict ahead + + # Intraday trading settings + TRADING_FREQUENCY = "5min" # '1min', '5min', '15min', '30min', '1h' + TRADING_HOURS = { + "start": "09:30", # Market open (ET) + "end": "16:00", # Market close (ET) + } + PRE_MARKET_HOURS = { + "start": "04:00", # Pre-market start + "end": "09:30", # Pre-market end + } + AFTER_HOURS = { + "start": "16:00", # After-hours start + "end": "20:00", # After-hours end + } + MAX_POSITION_HOLD_TIME = "4h" # Maximum time to hold a position + MIN_POSITION_HOLD_TIME = "10min" # Minimum time to hold a position + MAX_DAILY_POSITIONS = 50 # Maximum number of positions per day + MAX_POSITION_SIZE = 0.03 # Maximum % of portfolio per position (3%) + + # Data pipeline settings + LOOKBACK_WINDOW = 60 # Days for feature calculation + REALTIME_FEATURE_WINDOW = 30 # Number of data points for real-time features + DATA_BUFFER_SIZE = 5000 # Number of data points to keep in memory + DATA_FLUSH_INTERVAL = 300 # seconds - how often to flush data to database + + # Alternative data settings + NEWS_API_KEY = "your_news_api_key" + TWITTER_BEARER_TOKEN = "your_twitter_bearer_token" + REDDIT_CLIENT_ID = "your_reddit_client_id" + REDDIT_CLIENT_SECRET = "your_reddit_client_secret" + NEWS_LOOKBACK_DAYS = 7 # Number of days to look back for news + SOCIAL_MEDIA_LOOKBACK_DAYS = 3 # Number of days to look back for social media + + # Feature definitions (referenced by models) + NEWS_FEATURES = [ + "sentiment", + "volume", + "recency", + "source_reliability", + "topic_relevance", + ] + SOCIAL_FEATURES = [ + "twitter_sentiment", + "twitter_volume", + "reddit_sentiment", + "reddit_volume", + "social_momentum", + ] + INTRADAY_FEATURES = [ + "return", + "volatility", + "momentum", + "volume_momentum", + "bid_ask_spread", + "bid_ask_spread_pct", + "volume_imbalance", + "order_flow", + "vwap_deviation", + ] + + # Stateful prediction + STATEFUL_PREDICTION = False + REALTIME_UPDATE_INTERVAL = 60 # seconds + + # Backtesting settings + INITIAL_CAPITAL = 100000 + TRANSACTION_COST = 0.0005 # 0.05% per trade + SLIPPAGE_MODEL = "volume_curve" # 'volume_curve', 'constant', or 'none' + SLIPPAGE_RATE = 0.0002 # 0.02% slippage + + # Live trading settings + LIVE_DATA_ENABLED = True + DATA_PROVIDER = "polygon" # 'polygon', 'alphavantage', 'ib', 'tdameritrade' + POLYGON_API_KEY = "your_polygon_api_key" + ALPHA_VANTAGE_API_KEY = "your_alpha_vantage_api_key" + IB_HOST = "127.0.0.1" + IB_PORT = 7497 + IB_CLIENT_ID = 1 + + # WebSocket settings + WEBSOCKET_RECONNECT_DELAY = 5 # seconds + WEBSOCKET_MAX_RETRIES = 20 + WEBSOCKET_PING_INTERVAL = 30 # seconds + + # Data loading settings + NUM_WORKERS = 8 # Number of data loading workers + PREFETCH_FACTOR = 4 # Number of batches to prefetch + + # Online learning settings + ONLINE_LEARNING = True # Enable online learning + ONLINE_LEARNING_RATE = 0.0001 # Learning rate for online updates + ONLINE_LEARNING_INTERVAL = 3600 # seconds - how often to perform online learning + + # Risk management settings + MAX_DAILY_LOSS = 0.01 # 1% max daily loss + MAX_DRAWDOWN = 0.05 # 5% max drawdown + VOLATILITY_TARGET = 0.15 # Annualized volatility target + POSITION_SIZING = ( + "volatility_target" # 'volatility_target', 'equal_weight', 'kelly' + ) + + # Execution settings + EXECUTION_ALGORITHM = "vwap" # 'vwap', 'twap', 'pov', 'implementation_shortfall' + EXECUTION_TIME_HORIZON = "5min" # Time to complete execution + MARKET_IMPACT_MODEL = "kyle" # 'kyle', 'almgren_chriss', or 'none' + + +config = Config() diff --git a/data/external/.gitkeep b/data/external/.gitkeep new file mode 100644 index 0000000..e69de29 diff --git a/data/processed/.gitkeep b/data/processed/.gitkeep new file mode 100644 index 0000000..e69de29 diff --git a/data/raw/.gitkeep b/data/raw/.gitkeep new file mode 100644 index 0000000..e69de29 diff --git a/live_trading.py b/live_trading.py new file mode 100644 index 0000000..0d82106 --- /dev/null +++ b/live_trading.py @@ -0,0 +1,569 @@ +import asyncio +import logging +import time +from datetime import datetime, timedelta +from typing import Dict, List + +import torch +from src.trading.paper_broker import PaperTradingBroker +from src.trading.real_time_trader import RealTimeTrader + +from config import config +from src.amd.optimizations import AMDOptimizer +from src.data.live_data import LiveDataService +from src.data.pipeline import StockDataPipeline +from src.models.intraday_gnn import IntradayGNN +from src.models.trainer import GNNTrainer +from src.utils.memory_manager import MemoryManager + +# Configure logging +logging.basicConfig( + level=logging.INFO, + format="%(asctime)s - %(name)s - %(levelname)s - %(message)s", + handlers=[ + logging.FileHandler("live_trading_r9700.log"), + logging.StreamHandler(), + ], +) +logger = logging.getLogger(__name__) + + +class LiveTradingSystem: + def __init__(self): + # Initialize memory manager + self.memory_manager = MemoryManager() + logger.info(self.memory_manager.get_memory_stats()) + + # Initialize AMD optimizer + self.amd_optimizer = AMDOptimizer() + + # Initialize components + self.pipeline = StockDataPipeline() + self.broker = PaperTradingBroker(initial_cash=config.INITIAL_CAPITAL) + + # Load model with memory management + self._initialize_model() + + # Initialize trader + self.trader = RealTimeTrader(self.model, self.pipeline, self.broker) + + # Initialize live data service + self.live_data_service = LiveDataService(self.pipeline, self._on_market_data) + + # Trading state + self.last_trade_time = None + self.last_model_update_time = time.time() + + def _initialize_model(self): + """Initialize the GNN model with AMD optimizations""" + # Check memory before loading model + if not self.memory_manager.ensure_memory(8 * 1024**3): # 8GB + logger.warning("Not enough GPU memory for model. Falling back to CPU.") + config.DEVICE = "cpu" + + # In a real implementation, you would load a pre-trained model + # For this example, we'll create a new model with the right number of features + + # Get number of features (this would be determined from your actual data) + # For this example, we'll use a reasonable estimate + num_features = len(config.INTRADAY_FEATURES) + 5 # +5 for price features + + self.model = IntradayGNN(num_features, config.SEQUENCE_LENGTH) + self.trainer = GNNTrainer(self.model) + + # Load pre-trained weights if available + try: + self.trainer.load_model() + logger.info("Loaded pre-trained model") + except Exception as e: + logger.warning(f"Could not load pre-trained model: {str(e)}") + logger.info("Using randomly initialized model") + + # Optimize model for AMD GPU + self.model = self.amd_optimizer.optimize_model(self.model) + + # Log memory after loading model + self.memory_manager.log_memory_usage("[After Model Load]") + + async def _on_market_data( + self, ticker: str, timestamp: str, data_type: str, data: Dict + ): + """Callback for market data updates with memory management""" + # Check memory before processing + if not self.memory_manager.ensure_memory(500 * 1024**2): # 500MB + logger.warning("Skipping data processing due to memory constraints") + return + + # Check if we should generate trading signals + if not self._should_generate_signals(timestamp): + return + + # Get current data for all tickers + current_data = self._prepare_current_data(timestamp) + + if not current_data: + return + + # Generate trading signals + await self._generate_trading_signals(current_data, timestamp) + + # Periodically update the model with online learning + if ( + config.ONLINE_LEARNING + and time.time() - self.last_model_update_time + > config.ONLINE_LEARNING_INTERVAL + ): + await self._online_learning(timestamp) + self.last_model_update_time = time.time() + + # Memory management + self.memory_manager.auto_manage_memory(threshold=0.8) + + def _should_generate_signals(self, timestamp: str) -> bool: + """Determine if we should generate trading signals""" + current_time = datetime.strptime(timestamp, "%Y-%m-%d %H:%M:%S").time() + + # Check if market is open + market_open = datetime.strptime(config.TRADING_HOURS["start"], "%H:%M").time() + market_close = datetime.strptime(config.TRADING_HOURS["end"], "%H:%M").time() + + if not (market_open <= current_time <= market_close): + return False + + # Check if it's time to generate signals based on trading frequency + if config.TRADING_FREQUENCY == "1min": + return True + elif config.TRADING_FREQUENCY == "5min": + return current_time.minute % 5 == 0 + elif config.TRADING_FREQUENCY == "15min": + return current_time.minute % 15 == 0 + else: # Default to 1 minute + return True + + def _prepare_current_data(self, timestamp: str) -> Dict: + """Prepare current data for prediction with memory management""" + from torch_geometric.data import Data + + current_date = timestamp.split(" ")[0] + tickers = config.INITIAL_TICKERS + + # Get sequence data for each ticker + sequence_features = [] + valid_tickers = [] + + for ticker in tickers: + # Check memory before processing ticker + if not self.memory_manager.ensure_memory(50 * 1024**2): # 50MB + logger.warning(f"Skipping {ticker} due to memory constraints") + continue + + try: + # Get features for the ticker + features = self.pipeline.get_latest_features([ticker], timestamp) + + if ticker not in features: + continue + + # Get sequence of features + timestamps = self._generate_timestamps_for_date(current_date) + current_idx = ( + timestamps.index(timestamp) + if timestamp in timestamps + else len(timestamps) - 1 + ) + sequence_start = timestamps[ + max(0, current_idx - config.SEQUENCE_LENGTH + 1) + ] + + # Get features for the sequence + sequence_data = self._get_feature_sequence( + ticker, sequence_start, timestamp + ) + + if len(sequence_data) < config.SEQUENCE_LENGTH: + # Pad with zeros if sequence is too short + padding = np.zeros( + ( + config.SEQUENCE_LENGTH - len(sequence_data), + sequence_data.shape[1], + ) + ) + sequence_data = np.vstack([padding, sequence_data]) + + sequence_features.append(sequence_data) + valid_tickers.append(ticker) + + except Exception as e: + logger.warning(f"Error processing {ticker} for {timestamp}: {str(e)}") + continue + + if not sequence_features: + return None + + # Convert to tensor (num_stocks, sequence_length, num_features) + x = torch.tensor(np.array(sequence_features), dtype=torch.float32) + + # Create edges based on sector relationships + edge_index = [] + edge_weight = [] + + for i, ticker1 in enumerate(valid_tickers): + for j, ticker2 in enumerate(valid_tickers): + if i < j: + # Check memory before processing edge + if not self.memory_manager.ensure_memory(1 * 1024**2): # 1MB + logger.warning( + "Skipping edge creation due to memory constraints" + ) + continue + + try: + # Get sector relationship + pit1 = self.pipeline.get_point_in_time_data( + ticker1, datetime.strptime(current_date, "%Y-%m-%d") + ) + pit2 = self.pipeline.get_point_in_time_data( + ticker2, datetime.strptime(current_date, "%Y-%m-%d") + ) + + if ( + pit1["sector"] + and pit2["sector"] + and pit1["sector"] == pit2["sector"] + ): + # Calculate correlation of recent returns + lookback_start = datetime.strptime( + sequence_start, "%Y-%m-%d %H:%M:%S" + ) + lookback_end = datetime.strptime( + timestamp, "%Y-%m-%d %H:%M:%S" + ) + + returns1 = self._get_returns_sequence( + ticker1, lookback_start, lookback_end + ) + returns2 = self._get_returns_sequence( + ticker2, lookback_start, lookback_end + ) + + if len(returns1) > 5 and len(returns2) > 5: + corr = np.corrcoef(returns1, returns2)[0, 1] + if not np.isnan(corr): + edge_index.append([i, j]) + edge_weight.append(corr) + except Exception as e: + logger.warning( + f"Error creating edge between {ticker1} and {ticker2}: {str(e)}" + ) + continue + + # Convert to tensors + edge_index = ( + torch.tensor(edge_index, dtype=torch.long).t() + if edge_index + else torch.empty((2, 0), dtype=torch.long) + ) + edge_weight = ( + torch.tensor(edge_weight, dtype=torch.float32).unsqueeze(1) + if edge_weight + else torch.empty((0, 1), dtype=torch.float32) + ) + + # Create Data object + data = Data(x=x, edge_index=edge_index, edge_attr=edge_weight) + data.timestamp = timestamp + data.date = current_date + data.tickers = valid_tickers + + return data + + def _generate_timestamps_for_date(self, date: str) -> List[str]: + """Generate timestamps for a given date""" + # Get market open and close times + market_open = datetime.strptime(config.TRADING_HOURS["start"], "%H:%M").time() + market_close = datetime.strptime(config.TRADING_HOURS["end"], "%H:%M").time() + + # Create datetime objects for open and close + open_datetime = datetime.strptime(f"{date} {market_open}", "%Y-%m-%d %H:%M:%S") + close_datetime = datetime.strptime( + f"{date} {market_close}", "%Y-%m-%d %H:%M:%S" + ) + + # Generate timestamps based on trading frequency + if config.TRADING_FREQUENCY == "1min": + delta = timedelta(minutes=1) + elif config.TRADING_FREQUENCY == "5min": + delta = timedelta(minutes=5) + elif config.TRADING_FREQUENCY == "15min": + delta = timedelta(minutes=15) + elif config.TRADING_FREQUENCY == "30min": + delta = timedelta(minutes=30) + elif config.TRADING_FREQUENCY == "1h": + delta = timedelta(hours=1) + else: # Default to 1 minute + delta = timedelta(minutes=1) + + timestamps = [] + current = open_datetime + + while current <= close_datetime: + timestamps.append(current.strftime("%Y-%m-%d %H:%M:%S")) + current += delta + + return timestamps + + def _get_feature_sequence(self, ticker: str, start: str, end: str): + """Get feature sequence for a ticker between start and end timestamps""" + if ( + ticker not in self.pipeline.intraday_data + or "features" not in self.pipeline.intraday_data[ticker] + ): + return np.array([]) + + features = self.pipeline.intraday_data[ticker]["features"] + sequence = features.loc[start:end].values + + return sequence + + def _get_returns_sequence( + self, ticker: str, start: datetime, end: datetime + ) -> List[float]: + """Get returns sequence for a ticker between start and end times""" + if ( + ticker not in self.pipeline.intraday_data + or "price_bars" not in self.pipeline.intraday_data[ticker] + ): + return [] + + price_bars = self.pipeline.intraday_data[ticker]["price_bars"] + sequence = price_bars.loc[ + start.strftime("%Y-%m-%d %H:%M:%S") : end.strftime("%Y-%m-%d %H:%M:%S") + ] + + if len(sequence) < 2: + return [] + + returns = sequence["close"].pct_change().dropna().values + return returns.tolist() + + async def _generate_trading_signals(self, data: Dict, timestamp: str): + """Generate trading signals based on model predictions with memory management""" + if not data: + return + + # Check memory before prediction + if not self.memory_manager.ensure_memory(2 * 1024**3): # 2GB + logger.warning("Skipping prediction due to memory constraints") + return + + try: + # Get predictions from model + with torch.no_grad(): + data = data.to(config.DEVICE) + predictions = self.model(data).squeeze().cpu().numpy() + + # Generate signals for each ticker + for i, ticker in enumerate(data.tickers): + prediction = predictions[i] + + # Get current price + current_price = None + latest_data = self.live_data_service.get_latest_data(ticker) + if "price_bar" in latest_data: + current_price = latest_data["price_bar"]["close"] + + if not current_price: + continue + + # Skip if we already have a position in this stock + if ticker in self.trader.current_positions: + continue + + # Check holding period constraints + if self.trader._check_holding_period(ticker, timestamp): + continue + + # Generate signals based on prediction + if prediction > 0.002: # Buy signal + # Calculate position size + position_size = self.trader._calculate_position_size( + ticker, current_price + ) + + if position_size > 0: + # Create buy order + order = { + "ticker": ticker, + "action": "buy", + "quantity": position_size, + "price": current_price, + "timestamp": timestamp, + "type": "market", + } + + # Submit order + order_id = self.broker.submit_order(order) + if order_id: + self.trader.pending_orders[order_id] = order + logger.info( + f"Submitted buy order for {position_size} shares of {ticker} at {current_price}" + ) + + elif ( + prediction < -0.002 and ticker in self.trader.current_positions + ): # Sell signal + # Create sell order + order = { + "ticker": ticker, + "action": "sell", + "quantity": self.trader.current_positions[ticker], + "price": current_price, + "timestamp": timestamp, + "type": "market", + } + + # Submit order + order_id = self.broker.submit_order(order) + if order_id: + self.trader.pending_orders[order_id] = order + logger.info( + f"Submitted sell order for {self.trader.current_positions[ticker]} shares of {ticker} at {current_price}" + ) + + except Exception as e: + logger.error(f"Error generating trading signals: {str(e)}", exc_info=True) + self.memory_manager.empty_cache() + + async def _online_learning(self, timestamp: str): + """Perform online learning with new data and memory management""" + current_date = timestamp.split(" ")[0] + tickers = config.INITIAL_TICKERS + + # Create dataset for online learning + dataset = self.pipeline.create_intraday_dataset( + tickers, current_date, current_date + ) + + if not dataset: + return + + # Check memory before online learning + if not self.memory_manager.ensure_memory(4 * 1024**3): # 4GB + logger.warning("Skipping online learning due to memory constraints") + return + + try: + # Get the most recent data point + recent_data = dataset[-1] + + # Perform online update + loss = self.trainer.online_update(recent_data) + if loss is not None: + logger.info(f"Online learning update - Loss: {loss:.6f}") + + except Exception as e: + logger.error(f"Error during online learning: {str(e)}", exc_info=True) + self.memory_manager.empty_cache() + + async def run(self): + """Run the live trading system with AMD optimizations""" + logger.info( + "Starting live trading system with AMD Radeon R9700 AI Pro optimizations" + ) + + # Start live data service + data_task = asyncio.create_task(self.live_data_service.start()) + + # Start trader + trader_task = asyncio.create_task(self._run_trader()) + + # Start memory monitor + memory_task = asyncio.create_task(self._monitor_memory()) + + # Wait for tasks to complete + await asyncio.gather(data_task, trader_task, memory_task) + + async def _run_trader(self): + """Run the trader component with memory management""" + while True: + try: + # Check memory before processing + if not self.memory_manager.ensure_memory(1 * 1024**3): # 1GB + logger.warning( + "Skipping trader processing due to memory constraints" + ) + await asyncio.sleep(5) + continue + + # Process pending orders + self.trader._process_pending_orders() + + # Update portfolio value + self.trader._update_portfolio_value() + + # Check risk limits + if self.trader._check_risk_limits(): + logger.info("Risk limits exceeded. Stopping trading for the day.") + self.trader._close_all_positions() + + # Sleep for a short interval + await asyncio.sleep(1) + + except Exception as e: + logger.error(f"Error in trader: {str(e)}", exc_info=True) + self.memory_manager.empty_cache() + await asyncio.sleep(5) + + async def _monitor_memory(self): + """Monitor GPU memory usage""" + while True: + try: + memory_info = self.memory_manager.check_memory() + if memory_info["usage_percent"] > 85: + logger.info( + f"High GPU memory usage: {memory_info['usage_percent']:.2f}%. Clearing cache." + ) + self.memory_manager.empty_cache() + + await asyncio.sleep(60) # Check every minute + + except Exception as e: + logger.error(f"Error in memory monitor: {str(e)}") + await asyncio.sleep(60) + + async def stop(self): + """Stop the live trading system""" + logger.info("Stopping live trading system") + + # Stop live data service + await self.live_data_service.stop() + + # Stop trader + self.trader._close_all_positions() + + # Clear memory + self.memory_manager.empty_cache() + + logger.info("Live trading system stopped") + logger.info(self.memory_manager.get_memory_stats()) + + +async def main(): + # Initialize live trading system + trading_system = LiveTradingSystem() + + try: + # Run the system + await trading_system.run() + except KeyboardInterrupt: + logger.info("Received keyboard interrupt. Shutting down...") + except Exception as e: + logger.error(f"Error in live trading system: {str(e)}", exc_info=True) + finally: + # Clean up + await trading_system.stop() + logger.info("Live trading system shutdown complete") + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/main.py b/main.py new file mode 100644 index 0000000..096cf84 --- /dev/null +++ b/main.py @@ -0,0 +1,156 @@ +import logging +from datetime import datetime + +import matplotlib.pyplot as plt +import pandas as pd +from src.evaluation.backtester import GNNBacktester +from src.evaluation.metrics import calculate_performance_metrics, compare_to_benchmark +from src.utils.visualization import ( + plot_feature_importance, + plot_performance, + plot_trade_log, +) + +from config import config +from src.amd.optimizations import AMDOptimizer +from src.data.pipeline import StockDataPipeline +from src.models.gnn_model import CorporateActionAwareGNN +from src.models.trainer import GNNTrainer +from src.utils.memory_manager import MemoryManager + +# Configure logging +logging.basicConfig( + level=logging.INFO, + format="%(asctime)s - %(name)s - %(levelname)s - %(message)s", + handlers=[ + logging.FileHandler("stock_gnn_r9700.log"), + logging.StreamHandler(), + ], +) +logger = logging.getLogger(__name__) + + +def main(): + # Initialize memory manager + memory_manager = MemoryManager() + logger.info(memory_manager.get_memory_stats()) + + # Initialize AMD optimizer + amd_optimizer = AMDOptimizer() + + # Initialize data pipeline + logger.info("Initializing data pipeline") + pipeline = StockDataPipeline() + + # Update all data + logger.info("Updating all data sources") + pipeline.update_all_data() + + # Create datasets + logger.info("Creating training and validation datasets") + train_dataset = pipeline.create_training_dataset() + val_dataset = pipeline.create_validation_dataset( + start_date=config.TRAIN_END_DATE, end_date=config.VAL_END_DATE + ) + + logger.info(f"Training dataset size: {len(train_dataset)}") + logger.info(f"Validation dataset size: {len(val_dataset)}") + + # Initialize model + logger.info("Initializing GNN model") + # Get number of features from first data point + num_features = train_dataset[0].x.shape[1] + model = CorporateActionAwareGNN(num_features) + + # Optimize model for AMD GPU + model = amd_optimizer.optimize_model(model) + + # Train model + logger.info("Training model with AMD optimizations") + trainer = GNNTrainer(model) + train_losses, val_losses = trainer.train(train_dataset, val_dataset) + + # Plot training curves + plt.figure(figsize=(10, 5)) + plt.plot(train_losses, label="Training Loss") + plt.plot(val_losses, label="Validation Loss") + plt.title("Training and Validation Loss") + plt.xlabel("Epoch") + plt.ylabel("Loss") + plt.legend() + plt.savefig("training_curves.png") + plt.close() + + # Benchmark model + logger.info("Benchmarking model performance") + sample_data = train_dataset[0].to(config.DEVICE) + benchmark_results = trainer.benchmark(sample_data) + logger.info(f"Benchmark Results: {benchmark_results}") + + # Load best model + trainer.load_model() + + # Run backtest on validation set + logger.info("Running backtest on validation set") + backtester = GNNBacktester(model, pipeline) + portfolio_values, trade_log = backtester.run_backtest(val_dataset) + + # Get benchmark data + benchmark_data = pipeline.price_data[config.INDEX_TICKER] + benchmark_values = benchmark_data.loc[portfolio_values.index]["Adj Close"] + + # Calculate performance metrics + logger.info("Calculating performance metrics") + portfolio_returns = portfolio_values.pct_change().dropna() + benchmark_returns = benchmark_values.pct_change().dropna() + + metrics = calculate_performance_metrics(portfolio_returns, benchmark_returns) + comparison = compare_to_benchmark(portfolio_values, benchmark_values) + + # Print metrics + logger.info("\nPerformance Metrics:") + for metric, value in metrics.items(): + if isinstance(value, float): + logger.info(f"{metric.replace('_', ' ').title()}: {value:.4f}") + else: + logger.info(f"{metric.replace('_', ' ').title()}: {value}") + + logger.info("\nComparison to Benchmark:") + for metric, value in comparison.items(): + if isinstance(value, float): + logger.info(f"{metric.replace('_', ' ').title()}: {value:.4f}") + else: + logger.info(f"{metric.replace('_', ' ').title()}: {value}") + + # Plot performance + plot_performance(portfolio_values, benchmark_values, "portfolio_performance.png") + + # Plot trade log + plot_trade_log(trade_log, "trade_log.png") + + # Save results + results = { + "portfolio_values": portfolio_values, + "benchmark_values": benchmark_values, + "trade_log": trade_log, + "metrics": metrics, + "comparison": comparison, + } + + results_df = pd.DataFrame( + { + "date": portfolio_values.index, + "portfolio_value": portfolio_values.values, + "benchmark_value": benchmark_values.values, + } + ) + results_df.to_csv("backtest_results.csv", index=False) + + # Memory cleanup + memory_manager.empty_cache() + logger.info("Backtest completed. Results saved to backtest_results.csv") + logger.info(memory_manager.get_memory_stats()) + + +if __name__ == "__main__": + main() diff --git a/models/.gitkeep b/models/.gitkeep new file mode 100644 index 0000000..e69de29 diff --git a/notebooks/.gitkeep b/notebooks/.gitkeep new file mode 100644 index 0000000..e69de29 diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..1739a39 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,58 @@ +# Core requirements +numpy>=1.24.0 +pandas>=2.0.0 +scipy>=1.10.0 +scikit-learn>=1.2.0 +tqdm>=4.65.0 +matplotlib>=3.7.0 +seaborn>=0.12.0 +python-dateutil>=2.8.0 +pytz>=2023.3 +requests>=2.28.0 +websockets>=11.0 +SQLAlchemy>=2.0.0 +sqlite3>=3.40.0 # Part of Python standard library + +# AMD GPU support for PyTorch (ROCm 5.6) +torch>=2.1.0 # ROCm-compatible version +torchvision>=0.16.0 # ROCm-compatible version +torchaudio>=2.1.0 # ROCm-compatible version +--index-url https://download.pytorch.org/whl/rocm5.6 + +# PyTorch Geometric with AMD support +torch-geometric>=2.4.0 +torch-scatter>=2.1.2 +torch-sparse>=0.6.18 +torch-cluster>=1.6.2 +torch-spline-conv>=1.2.2 +--find-links https://data.pyg.org/whl/torch-2.1.0+rocm5.6.html + +# ROCm libraries for performance +rocblas>=3.1.0 +hipblaslt>=0.6.0 +miopen-hip>=2.19.0 +rccl>=2.15.5 + +# Alternative data processing +yfinance>=0.2.20 +alpha_vantage>=2.3.1 +polygon-api-client>=1.12.0 +tweepy>=4.14.0 +praw>=7.7.0 +newspaper3k>=0.2.8 +transformers>=4.30.0 +sentencepiece>=0.1.99 + +# Interactive Brokers +ib_insync>=0.9.86 + +# For AMD-specific optimizations +tensorboard>=2.13.0 +psutil>=5.9.0 + +# Web frontend +fastapi>=0.104.0 +uvicorn[standard]>=0.24.0 +jinja2>=3.1.0 +python-multipart>=0.0.6 +aiofiles>=23.2.0 diff --git a/src/__init__.py b/src/__init__.py new file mode 100644 index 0000000..521670b --- /dev/null +++ b/src/__init__.py @@ -0,0 +1 @@ +# src package diff --git a/src/amd/__init__.py b/src/amd/__init__.py new file mode 100644 index 0000000..dbb64ea --- /dev/null +++ b/src/amd/__init__.py @@ -0,0 +1 @@ +# AMD optimizations package diff --git a/src/amd/optimizations.py b/src/amd/optimizations.py new file mode 100644 index 0000000..d6a0690 --- /dev/null +++ b/src/amd/optimizations.py @@ -0,0 +1,349 @@ +import logging +import time + +import torch +import torch.nn as nn +import torch.nn.functional as F +from torch_geometric.nn import MessagePassing +from torch_geometric.utils import softmax + +from config import config + +logger = logging.getLogger(__name__) + + +class AMDOptimizer: + """ + AMD-specific optimizations for PyTorch models + """ + + def __init__(self): + self.device = torch.device(config.DEVICE) + self._configure_rocm() + + def _configure_rocm(self): + """Configure ROCm for optimal performance""" + if config.DEVICE == "cuda" and config.AMD_GPU: + try: + # Set ROCm optimization level + torch.backends.hip.set_optimization_level(config.ROCM_OPT_LEVEL) + + # Enable memory efficient attention if available + try: + from xformers.ops import MemoryEfficientAttentionFlashAttentionOp + + torch.backends.cuda.enable_flash_sdp(True) + logger.info("Enabled Flash Attention for AMD GPU") + except ImportError: + logger.warning("xformers not available, using standard attention") + + # Configure memory limits + total_memory = torch.cuda.get_device_properties(0).total_memory + memory_limit = int(total_memory * config.GPU_MEMORY_LIMIT) + torch.cuda.set_per_process_memory_fraction(config.GPU_MEMORY_LIMIT, 0) + + logger.info( + f"Configured ROCm with optimization level {config.ROCM_OPT_LEVEL}" + ) + logger.info( + f"GPU Memory: {total_memory / 1024**3:.2f}GB, Limit: {memory_limit / 1024**3:.2f}GB" + ) + + except Exception as e: + logger.error(f"Error configuring ROCm: {str(e)}") + + def optimize_model(self, model: nn.Module): + """Apply AMD-specific optimizations to a model""" + if config.DEVICE != "cuda" or not config.AMD_GPU: + return model + + try: + # Move model to GPU + model = model.to(self.device) + + # Apply mixed precision if enabled + if config.MIXED_PRECISION: + model = self._apply_mixed_precision(model) + + # Apply memory optimizations + model = self._apply_memory_optimizations(model) + + logger.info("Applied AMD optimizations to model") + return model + + except Exception as e: + logger.error(f"Error optimizing model: {str(e)}") + return model.to(self.device) + + def _apply_mixed_precision(self, model: nn.Module): + """Apply mixed precision training to the model""" + # Convert model to use mixed precision + if config.PRECISION == "fp16": + model = model.half() + elif config.PRECISION == "bf16": + model = model.to(torch.bfloat16) + + # Convert specific layers to full precision if needed + for name, module in model.named_modules(): + if isinstance(module, (nn.LayerNorm, nn.BatchNorm2d)): + module = module.float() + + logger.info(f"Applied mixed precision training with {config.PRECISION}") + return model + + def _apply_memory_optimizations(self, model: nn.Module): + """Apply memory optimizations to the model""" + # Enable gradient checkpointing for memory efficiency + if ( + hasattr(model, "supports_gradient_checkpointing") + and model.supports_gradient_checkpointing + ): + model.gradient_checkpointing_enable() + logger.info("Enabled gradient checkpointing") + + # Apply activation checkpointing to specific modules + for name, module in model.named_modules(): + if isinstance(module, (nn.LSTM, nn.GRU)): + module.activation_checkpointing = True + + return model + + def get_precision_dtype(self): + """Get the precision dtype for mixed precision training""" + if config.PRECISION == "fp16": + return torch.float16 + elif config.PRECISION == "bf16": + return torch.bfloat16 + else: + return torch.float32 + + def benchmark_model(self, model: nn.Module, input_data, num_runs: int = 100): + """Benchmark model performance on AMD GPU""" + if config.DEVICE != "cuda": + logger.warning("Benchmarking only supported on GPU") + return {} + + try: + # Warm up + for _ in range(10): + _ = model(input_data) + + # Benchmark inference + start_time = time.time() + for _ in range(num_runs): + with torch.no_grad(): + _ = model(input_data) + inference_time = (time.time() - start_time) / num_runs + + # Benchmark training + model.train() + optimizer = torch.optim.Adam(model.parameters(), lr=config.LEARNING_RATE) + criterion = nn.MSELoss() + + start_time = time.time() + for _ in range(num_runs): + optimizer.zero_grad() + out = model(input_data) + loss = criterion(out, torch.randn_like(out)) + loss.backward() + optimizer.step() + training_time = (time.time() - start_time) / num_runs + + # Memory usage + memory_allocated = torch.cuda.memory_allocated(0) + max_memory = torch.cuda.max_memory_allocated(0) + + return { + "inference_time": inference_time, + "training_time": training_time, + "throughput_inference": 1 / inference_time, + "throughput_training": 1 / training_time, + "memory_allocated": memory_allocated, + "max_memory": max_memory, + "memory_usage_percent": ( + memory_allocated / torch.cuda.get_device_properties(0).total_memory + ) + * 100, + } + + except Exception as e: + logger.error(f"Error benchmarking model: {str(e)}") + return {} + + +class AMDSparseAttention(nn.Module): + """ + Sparse attention implementation optimized for AMD GPUs + """ + + def __init__(self, embed_dim, num_heads, dropout=0.1): + super().__init__() + self.embed_dim = embed_dim + self.num_heads = num_heads + self.head_dim = embed_dim // num_heads + self.scaling = self.head_dim**-0.5 + + self.qkv_proj = nn.Linear(embed_dim, embed_dim * 3) + self.out_proj = nn.Linear(embed_dim, embed_dim) + self.dropout = nn.Dropout(dropout) + + # Initialize weights + self._init_weights() + + def _init_weights(self): + nn.init.xavier_uniform_(self.qkv_proj.weight) + nn.init.xavier_uniform_(self.out_proj.weight) + nn.init.zeros_(self.qkv_proj.bias) + nn.init.zeros_(self.out_proj.bias) + + def forward(self, x, mask=None): + batch_size, seq_len, embed_dim = x.size() + + # Project queries, keys, values + qkv = self.qkv_proj(x) + q, k, v = qkv.chunk(3, dim=-1) + + # Reshape for multi-head attention + q = q.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2) + k = k.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2) + v = v.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose(1, 2) + + # Compute attention scores + attn_scores = torch.matmul(q, k.transpose(-2, -1)) * self.scaling + + # Apply mask if provided + if mask is not None: + attn_scores = attn_scores.masked_fill(mask == 0, float("-inf")) + + # Compute attention weights + attn_weights = F.softmax(attn_scores, dim=-1) + attn_weights = self.dropout(attn_weights) + + # Apply attention to values + output = torch.matmul(attn_weights, v) + + # Concatenate heads + output = ( + output.transpose(1, 2).contiguous().view(batch_size, seq_len, embed_dim) + ) + + # Final projection + output = self.out_proj(output) + + return output + + +class AMDGATConv(MessagePassing): + """ + GATConv implementation optimized for AMD GPUs + """ + + def __init__( + self, + in_channels, + out_channels, + heads=1, + concat=True, + dropout=0.6, + add_self_loops=True, + ): + super().__init__(aggr="add", node_dim=0) + self.in_channels = in_channels + self.out_channels = out_channels + self.heads = heads + self.concat = concat + self.dropout = dropout + self.add_self_loops = add_self_loops + + # Linear transformations for each head + self.lin_src = nn.Parameter(torch.Tensor(in_channels, heads * out_channels)) + self.lin_dst = nn.Parameter(torch.Tensor(in_channels, heads * out_channels)) + + # Attention parameters + self.att_src = nn.Parameter(torch.Tensor(1, heads, out_channels)) + self.att_dst = nn.Parameter(torch.Tensor(1, heads, out_channels)) + + # Bias + self.bias = nn.Parameter(torch.Tensor(heads * out_channels)) + + # Initialize weights + self.reset_parameters() + + def reset_parameters(self): + nn.init.xavier_uniform_(self.lin_src) + nn.init.xavier_uniform_(self.lin_dst) + nn.init.xavier_uniform_(self.att_src) + nn.init.xavier_uniform_(self.att_dst) + nn.init.zeros_(self.bias) + + def forward(self, x, edge_index, edge_attr=None, size=None): + # Linear transformation + if size is None and torch.is_tensor(x): + x_src = x_dst = torch.matmul(x, self.lin_src).view( + -1, self.heads, self.out_channels + ) + else: + x_src, x_dst = x[0], x[1] + x_src = torch.matmul(x_src, self.lin_src).view( + -1, self.heads, self.out_channels + ) + x_dst = torch.matmul(x_dst, self.lin_dst).view( + -1, self.heads, self.out_channels + ) + + # Add self loops if needed + if self.add_self_loops: + num_nodes = x_src.size(0) + edge_index, edge_attr = self._add_self_loops( + edge_index, edge_attr, num_nodes + ) + + # Compute attention coefficients + alpha_src = (x_src * self.att_src).sum(dim=-1) + alpha_dst = (x_dst * self.att_dst).sum(dim=-1) + alpha = (alpha_src, alpha_dst) + + # Propagate + out = self.propagate(edge_index, x=(x_src, x_dst), alpha=alpha, size=size) + + # Concatenate or average heads + if self.concat: + out = out.view(-1, self.heads * self.out_channels) + else: + out = out.mean(dim=1) + + # Add bias + out = out + self.bias + + return out + + def _add_self_loops(self, edge_index, edge_attr, num_nodes): + # Add self loops to edge_index + loop_index = torch.arange( + 0, num_nodes, dtype=torch.long, device=edge_index.device + ) + loop_index = loop_index.unsqueeze(0).repeat(2, 1) + + if edge_attr is not None: + loop_attr = edge_attr.new_zeros((num_nodes,) + edge_attr.size()[1:]) + edge_attr = torch.cat([edge_attr, loop_attr], dim=0) + + edge_index = torch.cat([edge_index, loop_index], dim=1) + return edge_index, edge_attr + + def propagate(self, edge_index, size=None, **kwargs): + return super().propagate(edge_index, size=size, **kwargs) + + def message(self, x_j, alpha_j, alpha_i, index, ptr, size_i): + # Compute attention weights + alpha = alpha_j + alpha_i + alpha = F.leaky_relu(alpha, negative_slope=0.2) + alpha = self._softmax(alpha, index, ptr, size_i) + alpha = F.dropout(alpha, p=self.dropout, training=self.training) + + # Weighted sum of values + return x_j * alpha.unsqueeze(-1) + + def _softmax(self, src, index, ptr, num_nodes): + # Memory-efficient softmax + return softmax(src, index, ptr, num_nodes) diff --git a/src/data/__init__.py b/src/data/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/data/live_data.py b/src/data/live_data.py new file mode 100644 index 0000000..b799e4b --- /dev/null +++ b/src/data/live_data.py @@ -0,0 +1,908 @@ +import asyncio +import json +import logging +import time +from datetime import datetime, timedelta +from typing import Callable, Dict, List, Optional + +import numpy as np +import pandas as pd +import torch +import websockets + +from config import config +from src.data.pipeline import StockDataPipeline +from src.utils.helpers import generate_intraday_timestamps +from src.utils.memory_manager import MemoryManager + +logger = logging.getLogger(__name__) + + +class LiveDataService: + def __init__( + self, pipeline: StockDataPipeline, on_data_callback: Optional[Callable] = None + ): + self.pipeline = pipeline + self.on_data_callback = on_data_callback + self.memory_manager = MemoryManager() + + # Data buffers + self.data_buffer = {} # {ticker: {timestamp: data}} + self.price_bars = {} # {ticker: DataFrame} + self.order_book = {} # {ticker: DataFrame} + self.trades = {} # {ticker: DataFrame} + self.features = {} # {ticker: DataFrame} + + # Initialize data buffers + for ticker in config.INITIAL_TICKERS: + self.data_buffer[ticker] = {} + self.price_bars[ticker] = pd.DataFrame() + self.order_book[ticker] = pd.DataFrame() + self.trades[ticker] = pd.DataFrame() + self.features[ticker] = pd.DataFrame() + + # WebSocket connection + self.websocket = None + self.running = False + self.reconnect_attempts = 0 + self.last_flush_time = time.time() + self.last_feature_update_time = time.time() + + # Subscriptions + self.subscribed_tickers = set(config.INITIAL_TICKERS) + + # Initialize database connection pool + self.db_pool = None + self._init_db_pool() + + def _init_db_pool(self): + """Initialize database connection pool for better performance""" + try: + import sqlite3 + + self.db_pool = sqlite3.connect( + self.pipeline.db_path, check_same_thread=False + ) + self.db_pool.execute("PRAGMA journal_mode=WAL") + self.db_pool.execute("PRAGMA cache_size=-10000") # 10MB cache + logger.info("Initialized database connection pool") + except Exception as e: + logger.error(f"Error initializing database pool: {str(e)}") + self.db_pool = None + + async def start(self): + """Start the live data service with AMD optimizations""" + self.running = True + logger.info("Starting live data service with AMD optimizations") + + # Start memory monitor + self.memory_manager.monitor_memory(interval=60) + + while self.running: + try: + if config.DATA_PROVIDER == "polygon": + await self._connect_polygon() + elif config.DATA_PROVIDER == "alphavantage": + await self._connect_alpha_vantage() + elif config.DATA_PROVIDER == "ib": + await self._connect_interactive_brokers() + else: + logger.error(f"Unsupported data provider: {config.DATA_PROVIDER}") + return + + except Exception as e: + logger.error(f"Error in live data service: {str(e)}", exc_info=True) + await self._handle_disconnect() + if self.reconnect_attempts < config.WEBSOCKET_MAX_RETRIES: + await asyncio.sleep(config.WEBSOCKET_RECONNECT_DELAY) + else: + logger.error( + "Max reconnection attempts reached. Stopping live data service." + ) + self.running = False + + async def stop(self): + """Stop the live data service""" + self.running = False + if self.websocket: + await self.websocket.close() + if self.db_pool: + self.db_pool.close() + logger.info("Live data service stopped") + + async def _connect_polygon(self): + """Connect to Polygon.io WebSocket with AMD optimizations""" + uri = f"wss://socket.polygon.io/stocks" + + async with websockets.connect( + uri, ping_interval=config.WEBSOCKET_PING_INTERVAL + ) as websocket: + self.websocket = websocket + self.reconnect_attempts = 0 + logger.info("Connected to Polygon.io WebSocket") + + # Authenticate + auth_msg = {"action": "auth", "params": config.POLYGON_API_KEY} + await websocket.send(json.dumps(auth_msg)) + + # Subscribe to tickers + await self._subscribe_polygon_tickers() + + # Start processing messages + async for message in websocket: + try: + # Check if we should stop + if not self.running: + break + + data = json.loads(message) + + # Handle different message types + if isinstance(data, list): + for msg in data: + await self._process_polygon_message(msg) + else: + await self._process_polygon_message(data) + + # Periodically flush data to database + if time.time() - self.last_flush_time > config.DATA_FLUSH_INTERVAL: + await self._flush_data_to_database() + self.last_flush_time = time.time() + + # Periodically update features + if ( + time.time() - self.last_feature_update_time + > config.REALTIME_UPDATE_INTERVAL + ): + await self._update_all_features() + self.last_feature_update_time = time.time() + + except Exception as e: + logger.error(f"Error processing message: {str(e)}", exc_info=True) + continue + + async def _subscribe_polygon_tickers(self): + """Subscribe to Polygon.io tickers with AMD optimizations""" + if not self.subscribed_tickers: + self.subscribed_tickers = set(config.INITIAL_TICKERS) + + # Subscribe to trades and quotes for each ticker + for ticker in self.subscribed_tickers: + subscribe_msg = { + "action": "subscribe", + "params": f"T.{ticker},Q.{ticker}", + } + await self.websocket.send(json.dumps(subscribe_msg)) + + logger.info( + f"Subscribed to {len(self.subscribed_tickers)} tickers on Polygon.io" + ) + + async def _process_polygon_message(self, msg: Dict): + """Process a message from Polygon.io with AMD optimizations""" + if msg.get("ev") == "T": # Trade message + await self._process_trade_message(msg) + elif msg.get("ev") == "Q": # Quote message + await self._process_quote_message(msg) + elif msg.get("ev") == "status": # Status message + logger.info(f"Polygon.io status: {msg.get('message')}") + + async def _process_trade_message(self, msg: Dict): + """Process a trade message with AMD optimizations""" + ticker = msg.get("sym") + if ticker not in self.subscribed_tickers: + return + + # Check memory before processing + if not self.memory_manager.ensure_memory(1 * 1024**2): # 1MB + logger.warning( + f"Skipping trade message for {ticker} due to memory constraints" + ) + return + + try: + timestamp = pd.to_datetime(msg.get("t"), unit="ms").strftime( + "%Y-%m-%d %H:%M:%S" + ) + price = msg.get("p") + size = msg.get("s") + conditions = msg.get("c", []) + + # Store trade data + if ticker not in self.data_buffer: + self.data_buffer[ticker] = {} + + if "trades" not in self.data_buffer[ticker]: + self.data_buffer[ticker]["trades"] = {} + + self.data_buffer[ticker]["trades"][timestamp] = { + "price": price, + "size": size, + "conditions": conditions, + } + + # Update price bars + await self._update_price_bars(ticker, timestamp, price, size) + + # Call callback if provided + if self.on_data_callback: + await self.on_data_callback( + ticker, + timestamp, + "trade", + {"price": price, "size": size, "conditions": conditions}, + ) + + except Exception as e: + logger.error( + f"Error processing trade message for {ticker}: {str(e)}", + exc_info=True, + ) + self.memory_manager.empty_cache() + + async def _process_quote_message(self, msg: Dict): + """Process a quote message with AMD optimizations""" + ticker = msg.get("sym") + if ticker not in self.subscribed_tickers: + return + + # Check memory before processing + if not self.memory_manager.ensure_memory(1 * 1024**2): # 1MB + logger.warning( + f"Skipping quote message for {ticker} due to memory constraints" + ) + return + + try: + timestamp = pd.to_datetime(msg.get("t"), unit="ms").strftime( + "%Y-%m-%d %H:%M:%S" + ) + bid_price = msg.get("bp") + bid_size = msg.get("bs") + ask_price = msg.get("ap") + ask_size = msg.get("as") + + # Store order book data + if ticker not in self.data_buffer: + self.data_buffer[ticker] = {} + + if "order_book" not in self.data_buffer[ticker]: + self.data_buffer[ticker]["order_book"] = {} + + self.data_buffer[ticker]["order_book"][timestamp] = { + "bid_price": bid_price, + "bid_size": bid_size, + "ask_price": ask_price, + "ask_size": ask_size, + } + + # Call callback if provided + if self.on_data_callback: + await self.on_data_callback( + ticker, + timestamp, + "quote", + { + "bid_price": bid_price, + "bid_size": bid_size, + "ask_price": ask_price, + "ask_size": ask_size, + }, + ) + + except Exception as e: + logger.error( + f"Error processing quote message for {ticker}: {str(e)}", + exc_info=True, + ) + self.memory_manager.empty_cache() + + async def _update_price_bars( + self, ticker: str, timestamp: str, price: float, size: int + ): + """Update price bars with new trade data with AMD optimizations""" + # Check memory before processing + if not self.memory_manager.ensure_memory(5 * 1024**2): # 5MB + logger.warning( + f"Skipping price bar update for {ticker} due to memory constraints" + ) + return + + try: + # Get current date + date = timestamp.split(" ")[0] + + # Determine the current bar based on trading frequency + current_time = datetime.strptime(timestamp, "%Y-%m-%d %H:%M:%S").time() + + if config.TRADING_FREQUENCY == "1min": + bar_time = current_time.replace(second=0, microsecond=0) + elif config.TRADING_FREQUENCY == "5min": + minute = (current_time.minute // 5) * 5 + bar_time = current_time.replace(minute=minute, second=0, microsecond=0) + elif config.TRADING_FREQUENCY == "15min": + minute = (current_time.minute // 15) * 15 + bar_time = current_time.replace(minute=minute, second=0, microsecond=0) + else: # Default to 1 minute + bar_time = current_time.replace(second=0, microsecond=0) + + bar_timestamp = f"{date} {bar_time.strftime('%H:%M:%S')}" + + # Update or create the current bar + if bar_timestamp in self.price_bars[ticker].index: + # Update existing bar + bar = self.price_bars[ticker].loc[bar_timestamp] + bar["high"] = max(bar["high"], price) + bar["low"] = min(bar["low"], price) + bar["close"] = price + bar["volume"] += size + bar["trades"] += 1 + + # Update VWAP + if "vwap" in bar: + bar["vwap"] = ( + (bar["vwap"] * (bar["volume"] - size)) + (price * size) + ) / bar["volume"] + else: + bar["vwap"] = price + + self.price_bars[ticker].loc[bar_timestamp] = bar + else: + # Create new bar + new_bar = { + "open": price, + "high": price, + "low": price, + "close": price, + "volume": size, + "vwap": price, + "trades": 1, + } + self.price_bars[ticker].loc[bar_timestamp] = new_bar + + # Keep only recent data to limit memory usage + if len(self.price_bars[ticker]) > config.DATA_BUFFER_SIZE: + self.price_bars[ticker] = self.price_bars[ticker].iloc[ + -config.DATA_BUFFER_SIZE : + ] + + except Exception as e: + logger.error( + f"Error updating price bars for {ticker}: {str(e)}", exc_info=True + ) + self.memory_manager.empty_cache() + + async def _update_all_features(self): + """Update features for all tickers with AMD optimizations""" + for ticker in self.subscribed_tickers: + await self._calculate_realtime_features(ticker) + + async def _calculate_realtime_features(self, ticker: str): + """Calculate real-time features for a ticker with AMD optimizations""" + # Check memory before processing + if not self.memory_manager.ensure_memory(10 * 1024**2): # 10MB + logger.warning( + f"Skipping feature calculation for {ticker} due to memory constraints" + ) + return + + try: + if ticker not in self.price_bars or self.price_bars[ticker].empty: + return + + # Get the most recent bars + recent_bars = self.price_bars[ticker].iloc[ + -config.REALTIME_FEATURE_WINDOW : + ] + + if len(recent_bars) < 5: # Need at least 5 bars for meaningful features + return + + # Calculate returns + recent_bars["returns"] = recent_bars["close"].pct_change() + + # Calculate volatility (annualized) + volatility = recent_bars["returns"].std() * np.sqrt(252) + + # Calculate momentum + momentum = recent_bars["returns"].mean() + + # Calculate volume momentum + volume_momentum = recent_bars["volume"].pct_change().mean() + + # Calculate bid-ask spread if order book data exists + bid_ask_spread = None + bid_ask_spread_pct = None + if not self.order_book[ticker].empty: + # Get most recent order book data + recent_order_book = self.order_book[ticker].iloc[-1] + bid_ask_spread = ( + recent_order_book["ask_price"] - recent_order_book["bid_price"] + ) + bid_ask_spread_pct = bid_ask_spread / ( + (recent_order_book["ask_price"] + recent_order_book["bid_price"]) + / 2 + ) + + # Calculate volume imbalance if order book data exists + volume_imbalance = None + if not self.order_book[ticker].empty: + recent_order_book = self.order_book[ticker].iloc[-1] + volume_imbalance = ( + recent_order_book["bid_size"] - recent_order_book["ask_size"] + ) / (recent_order_book["bid_size"] + recent_order_book["ask_size"]) + + # Calculate order flow if trade data exists + order_flow = None + if not self.trades[ticker].empty and not self.order_book[ticker].empty: + # Get recent trades and order book data for the same period + recent_trades = self.trades[ticker].iloc[ + -config.REALTIME_FEATURE_WINDOW : + ] + recent_order_book = self.order_book[ticker].iloc[ + -config.REALTIME_FEATURE_WINDOW : + ] + + # Merge trades with order book data + merged = pd.merge_asof( + recent_trades.sort_index(), + recent_order_book.sort_index(), + left_index=True, + right_index=True, + direction="backward", + ) + + # Classify trades as buyer or seller initiated + if not merged.empty: + merged["trade_sign"] = np.where( + merged["price"] + > (merged["bid_price"] + merged["ask_price"]) / 2, + 1, # Buyer-initiated + -1, # Seller-initiated + ) + + # Aggregate order flow + order_flow = merged["trade_sign"].sum() + + # Calculate VWAP deviation + vwap_deviation = ( + recent_bars["close"].iloc[-1] / recent_bars["vwap"].iloc[-1] - 1 + if "vwap" in recent_bars + else None + ) + + # Store features + timestamp = recent_bars.index[-1] + features = { + "timestamp": timestamp, + "returns": recent_bars["returns"].iloc[-1], + "volatility": volatility, + "momentum": momentum, + "volume_momentum": volume_momentum, + "bid_ask_spread": bid_ask_spread, + "bid_ask_spread_pct": bid_ask_spread_pct, + "volume_imbalance": volume_imbalance, + "order_flow": order_flow, + "vwap_deviation": vwap_deviation, + } + + # Add to features DataFrame + self.features[ticker].loc[timestamp] = features + + # Keep only recent features to limit memory usage + if len(self.features[ticker]) > config.DATA_BUFFER_SIZE: + self.features[ticker] = self.features[ticker].iloc[ + -config.DATA_BUFFER_SIZE : + ] + + except Exception as e: + logger.error( + f"Error calculating features for {ticker}: {str(e)}", exc_info=True + ) + self.memory_manager.empty_cache() + + async def _flush_data_to_database(self): + """Flush buffered data to the database with AMD optimizations""" + logger.info("Flushing data to database") + + # Check memory before flushing + if not self.memory_manager.ensure_memory(500 * 1024**2): # 500MB + logger.warning("Skipping data flush due to memory constraints") + return + + try: + for ticker in self.subscribed_tickers: + # Flush price bars + if not self.price_bars[ticker].empty: + await self._store_price_bars(ticker, self.price_bars[ticker]) + + # Flush order book data + if ( + ticker in self.data_buffer + and "order_book" in self.data_buffer[ticker] + and self.data_buffer[ticker]["order_book"] + ): + order_book_data = pd.DataFrame.from_dict( + self.data_buffer[ticker]["order_book"], orient="index" + ) + await self._store_order_book_data(ticker, order_book_data) + + # Flush trade data + if ( + ticker in self.data_buffer + and "trades" in self.data_buffer[ticker] + and self.data_buffer[ticker]["trades"] + ): + trade_data = pd.DataFrame.from_dict( + self.data_buffer[ticker]["trades"], orient="index" + ) + await self._store_trade_data(ticker, trade_data) + + # Flush features + if not self.features[ticker].empty: + await self._store_features(ticker, self.features[ticker]) + + # Clear buffer after flushing + self.data_buffer = {ticker: {} for ticker in self.subscribed_tickers} + + except Exception as e: + logger.error(f"Error flushing data to database: {str(e)}", exc_info=True) + self.memory_manager.empty_cache() + + async def _store_price_bars(self, ticker: str, price_bars: pd.DataFrame): + """Store price bars in the database with AMD optimizations""" + if price_bars.empty: + return + + # Convert to list of dictionaries for bulk insert + price_bars_list = [] + for timestamp, row in price_bars.iterrows(): + price_bars_list.append( + { + "ticker": ticker, + "timestamp": timestamp, + "open": row["open"], + "high": row["high"], + "low": row["low"], + "close": row["close"], + "volume": row["volume"], + "vwap": row.get("vwap", None), + "trades": row.get("trades", None), + } + ) + + # Use connection pool for better performance + if self.db_pool: + try: + cursor = self.db_pool.cursor() + cursor.executemany( + """ + INSERT OR REPLACE INTO price_bars + (ticker, timestamp, open, high, low, close, volume, vwap, trades) + VALUES (:ticker, :timestamp, :open, :high, :low, :close, :volume, :vwap, :trades) + """, + price_bars_list, + ) + self.db_pool.commit() + except Exception as e: + logger.error(f"Error storing price bars for {ticker}: {str(e)}") + self.db_pool.rollback() + else: + # Fallback to regular connection + with sqlite3.connect(self.pipeline.db_path) as conn: + cursor = conn.cursor() + cursor.executemany( + """ + INSERT OR REPLACE INTO price_bars + (ticker, timestamp, open, high, low, close, volume, vwap, trades) + VALUES (:ticker, :timestamp, :open, :high, :low, :close, :volume, :vwap, :trades) + """, + price_bars_list, + ) + conn.commit() + + async def _store_order_book_data(self, ticker: str, order_book: pd.DataFrame): + """Store order book data in the database with AMD optimizations""" + if order_book.empty: + return + + # Convert to list of dictionaries for bulk insert + order_book_list = [] + for timestamp, row in order_book.iterrows(): + order_book_list.append( + { + "ticker": ticker, + "timestamp": timestamp, + "bid_price": row["bid_price"], + "bid_size": row["bid_size"], + "ask_price": row["ask_price"], + "ask_size": row["ask_size"], + } + ) + + # Use connection pool for better performance + if self.db_pool: + try: + cursor = self.db_pool.cursor() + cursor.executemany( + """ + INSERT OR REPLACE INTO order_book + (ticker, timestamp, bid_price, bid_size, ask_price, ask_size) + VALUES (:ticker, :timestamp, :bid_price, :bid_size, :ask_price, :ask_size) + """, + order_book_list, + ) + self.db_pool.commit() + except Exception as e: + logger.error(f"Error storing order book data for {ticker}: {str(e)}") + self.db_pool.rollback() + else: + # Fallback to regular connection + with sqlite3.connect(self.pipeline.db_path) as conn: + cursor = conn.cursor() + cursor.executemany( + """ + INSERT OR REPLACE INTO order_book + (ticker, timestamp, bid_price, bid_size, ask_price, ask_size) + VALUES (:ticker, :timestamp, :bid_price, :bid_size, :ask_price, :ask_size) + """, + order_book_list, + ) + conn.commit() + + async def _store_trade_data(self, ticker: str, trades: pd.DataFrame): + """Store trade data in the database with AMD optimizations""" + if trades.empty: + return + + # Convert to list of dictionaries for bulk insert + trades_list = [] + for timestamp, row in trades.iterrows(): + trades_list.append( + { + "ticker": ticker, + "timestamp": timestamp, + "price": row["price"], + "size": row["size"], + "trade_condition": ( + row.get("conditions", [None])[0] + if isinstance(row.get("conditions"), list) + else None + ), + } + ) + + # Use connection pool for better performance + if self.db_pool: + try: + cursor = self.db_pool.cursor() + cursor.executemany( + """ + INSERT OR REPLACE INTO trades + (ticker, timestamp, price, size, trade_condition) + VALUES (:ticker, :timestamp, :price, :size, :trade_condition) + """, + trades_list, + ) + self.db_pool.commit() + except Exception as e: + logger.error(f"Error storing trade data for {ticker}: {str(e)}") + self.db_pool.rollback() + else: + # Fallback to regular connection + with sqlite3.connect(self.pipeline.db_path) as conn: + cursor = conn.cursor() + cursor.executemany( + """ + INSERT OR REPLACE INTO trades + (ticker, timestamp, price, size, trade_condition) + VALUES (:ticker, :timestamp, :price, :size, :trade_condition) + """, + trades_list, + ) + conn.commit() + + async def _store_features(self, ticker: str, features: pd.DataFrame): + """Store features in the database with AMD optimizations""" + if features.empty: + return + + # Convert to list of dictionaries for bulk insert + features_list = [] + for timestamp, row in features.iterrows(): + for feature_name, feature_value in row.items(): + features_list.append( + { + "ticker": ticker, + "timestamp": timestamp, + "feature_name": feature_name, + "feature_value": feature_value, + } + ) + + # Use connection pool for better performance + if self.db_pool: + try: + cursor = self.db_pool.cursor() + cursor.executemany( + """ + INSERT OR REPLACE INTO features + (ticker, timestamp, feature_name, feature_value) + VALUES (:ticker, :timestamp, :feature_name, :feature_value) + """, + features_list, + ) + self.db_pool.commit() + except Exception as e: + logger.error(f"Error storing features for {ticker}: {str(e)}") + self.db_pool.rollback() + else: + # Fallback to regular connection + with sqlite3.connect(self.pipeline.db_path) as conn: + cursor = conn.cursor() + cursor.executemany( + """ + INSERT OR REPLACE INTO features + (ticker, timestamp, feature_name, feature_value) + VALUES (:ticker, :timestamp, :feature_name, :feature_value) + """, + features_list, + ) + conn.commit() + + async def _handle_disconnect(self): + """Handle WebSocket disconnection with AMD optimizations""" + self.reconnect_attempts += 1 + logger.warning( + f"WebSocket disconnected. Attempt {self.reconnect_attempts} of {config.WEBSOCKET_MAX_RETRIES}" + ) + + if self.websocket: + await self.websocket.close() + self.websocket = None + + # Clear memory on disconnect + self.memory_manager.empty_cache() + + async def subscribe(self, tickers: List[str]): + """Subscribe to additional tickers with AMD optimizations""" + new_tickers = set(tickers) - self.subscribed_tickers + if not new_tickers: + return + + self.subscribed_tickers.update(new_tickers) + + # Initialize data structures for new tickers + for ticker in new_tickers: + self.data_buffer[ticker] = {} + self.price_bars[ticker] = pd.DataFrame() + self.order_book[ticker] = pd.DataFrame() + self.trades[ticker] = pd.DataFrame() + self.features[ticker] = pd.DataFrame() + + # Subscribe to new tickers + if config.DATA_PROVIDER == "polygon" and self.websocket: + for ticker in new_tickers: + subscribe_msg = { + "action": "subscribe", + "params": f"T.{ticker},Q.{ticker}", + } + await self.websocket.send(json.dumps(subscribe_msg)) + + logger.info( + f"Subscribed to {len(new_tickers)} new tickers: {', '.join(new_tickers)}" + ) + + async def unsubscribe(self, tickers: List[str]): + """Unsubscribe from tickers with AMD optimizations""" + removed_tickers = set(tickers) & self.subscribed_tickers + if not removed_tickers: + return + + self.subscribed_tickers -= removed_tickers + + # Remove from data structures + for ticker in removed_tickers: + if ticker in self.data_buffer: + del self.data_buffer[ticker] + if ticker in self.price_bars: + del self.price_bars[ticker] + if ticker in self.order_book: + del self.order_book[ticker] + if ticker in self.trades: + del self.trades[ticker] + if ticker in self.features: + del self.features[ticker] + + # Unsubscribe from tickers + if config.DATA_PROVIDER == "polygon" and self.websocket: + for ticker in removed_tickers: + unsubscribe_msg = { + "action": "unsubscribe", + "params": f"T.{ticker},Q.{ticker}", + } + await self.websocket.send(json.dumps(unsubscribe_msg)) + + logger.info( + f"Unsubscribed from {len(removed_tickers)} tickers: {', '.join(removed_tickers)}" + ) + + def get_latest_data(self, ticker: str) -> Dict: + """Get the latest data for a ticker with AMD optimizations""" + if ticker not in self.data_buffer: + return {} + + # Check memory before processing + if not self.memory_manager.ensure_memory(1 * 1024**2): # 1MB + logger.warning( + f"Skipping latest data retrieval for {ticker} due to memory constraints" + ) + return {} + + latest_data = {} + + # Get latest price bar + if not self.price_bars[ticker].empty: + latest_data["price_bar"] = self.price_bars[ticker].iloc[-1].to_dict() + + # Get latest order book data + if ( + "order_book" in self.data_buffer[ticker] + and self.data_buffer[ticker]["order_book"] + ): + latest_data["order_book"] = list( + self.data_buffer[ticker]["order_book"].values() + )[-1] + + # Get latest trade data + if "trades" in self.data_buffer[ticker] and self.data_buffer[ticker]["trades"]: + latest_data["trade"] = list(self.data_buffer[ticker]["trades"].values())[-1] + + # Get latest features + if not self.features[ticker].empty: + latest_data["features"] = self.features[ticker].iloc[-1].to_dict() + + return latest_data + + def get_latest_features(self, ticker: str) -> Dict: + """Get the latest features for a ticker with AMD optimizations""" + if ticker not in self.features or self.features[ticker].empty: + return {} + + # Check memory before processing + if not self.memory_manager.ensure_memory(1 * 1024**2): # 1MB + logger.warning( + f"Skipping latest features retrieval for {ticker} due to memory constraints" + ) + return {} + + return self.features[ticker].iloc[-1].to_dict() + + def get_latest_features_batch(self, tickers: List[str]) -> Dict[str, Dict]: + """Get the latest features for multiple tickers with AMD optimizations""" + features = {} + + # Check memory before processing + if not self.memory_manager.ensure_memory(10 * 1024**2): # 10MB + logger.warning( + "Skipping batch features retrieval due to memory constraints" + ) + return features + + for ticker in tickers: + if ticker in self.features and not self.features[ticker].empty: + features[ticker] = self.features[ticker].iloc[-1].to_dict() + + return features + + async def _connect_alpha_vantage(self): + """Placeholder for Alpha Vantage connection""" + logger.warning("Alpha Vantage live connection not yet implemented") + while self.running: + await asyncio.sleep(60) + + async def _connect_interactive_brokers(self): + """Placeholder for Interactive Brokers connection""" + logger.warning("Interactive Brokers live connection not yet implemented") + while self.running: + await asyncio.sleep(60) diff --git a/src/data/news_processor.py b/src/data/news_processor.py new file mode 100644 index 0000000..dc42885 --- /dev/null +++ b/src/data/news_processor.py @@ -0,0 +1,36 @@ +""" +News data processor for fetching and analyzing financial news. +""" + +import logging +from typing import Dict, List + +from config import config + +logger = logging.getLogger(__name__) + + +class NewsProcessor: + """Process news data for tickers.""" + + def __init__(self): + self.news_data = {} + + def fetch_news(self, tickers: List[str], start_date: str, end_date: str): + """Fetch news articles for the given tickers and date range.""" + logger.info( + f"Fetching news for {len(tickers)} tickers from {start_date} to {end_date}" + ) + # Placeholder: integrate with a news API (e.g., NewsAPI, Bloomberg) + pass + + def get_news_features(self, ticker: str, date: str) -> Dict: + """Return news-based features for a ticker on a specific date.""" + # Placeholder: return default features + return { + "news_sentiment": 0.0, + "news_volume": 0, + "news_recency": 0.0, + "news_source_reliability": 0.0, + "news_topic_relevance": 0.0, + } diff --git a/src/data/pipeline.py b/src/data/pipeline.py new file mode 100644 index 0000000..cbab3b6 --- /dev/null +++ b/src/data/pipeline.py @@ -0,0 +1,1131 @@ +import logging +import os +import pickle +import sqlite3 +from datetime import datetime, timedelta +from typing import Dict, List, Optional, Tuple + +import numpy as np +import pandas as pd +import torch +from tqdm import tqdm + +from config import config +from src.data.news_processor import NewsProcessor +from src.data.social_processor import SocialMediaProcessor +from src.utils.helpers import generate_intraday_timestamps +from src.utils.memory_manager import MemoryManager + +logger = logging.getLogger(__name__) + + +class StockDataPipeline: + def __init__(self): + self.db_path = os.path.join(config.PROCESSED_DATA_DIR, "stock_data.db") + self._init_db() + + # Initialize memory manager + self.memory_manager = MemoryManager() + + # Initialize components + self.price_data = {} + self.corporate_actions = {} + self.sector_data = {} + self.index_composition = {} + self.delisted_tickers = set() + self.intraday_data = {} + + # Initialize alternative data processors + self.news_processor = NewsProcessor() + self.social_processor = SocialMediaProcessor() + + # Load existing data if available + self._load_data() + + # Load delisted tickers + self._load_delisted_tickers() + + def _init_db(self): + """Initialize the SQLite database with AMD-optimized settings""" + with sqlite3.connect(self.db_path) as conn: + cursor = conn.cursor() + + # Enable WAL mode for better concurrency (important for AMD GPUs) + cursor.execute("PRAGMA journal_mode=WAL") + + # Increase cache size for better performance + cursor.execute("PRAGMA cache_size=-10000") # 10MB cache + + # Create tables for different data types + cursor.execute( + """ + CREATE TABLE IF NOT EXISTS price_data ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + ticker TEXT, + date TEXT, + open REAL, + high REAL, + low REAL, + close REAL, + adj_close REAL, + volume INTEGER, + UNIQUE(ticker, date) + ) + """ + ) + + cursor.execute( + """ + CREATE TABLE IF NOT EXISTS corporate_actions ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + ticker TEXT, + date TEXT, + action_type TEXT, + value REAL, + details TEXT, + UNIQUE(ticker, date, action_type) + ) + """ + ) + + cursor.execute( + """ + CREATE TABLE IF NOT EXISTS sector_data ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + ticker TEXT, + date TEXT, + sector TEXT, + industry TEXT, + UNIQUE(ticker, date) + ) + """ + ) + + cursor.execute( + """ + CREATE TABLE IF NOT EXISTS index_composition ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + index_ticker TEXT, + date TEXT, + member_ticker TEXT, + UNIQUE(index_ticker, date, member_ticker) + ) + """ + ) + + cursor.execute( + """ + CREATE TABLE IF NOT EXISTS features ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + ticker TEXT, + date TEXT, + feature_name TEXT, + feature_value REAL, + UNIQUE(ticker, date, feature_name) + ) + """ + ) + + # Create indexes for faster queries + cursor.execute( + "CREATE INDEX IF NOT EXISTS idx_price_data_ticker_date ON price_data(ticker, date)" + ) + cursor.execute( + "CREATE INDEX IF NOT EXISTS idx_corporate_actions_ticker_date ON corporate_actions(ticker, date)" + ) + cursor.execute( + "CREATE INDEX IF NOT EXISTS idx_sector_data_ticker_date ON sector_data(ticker, date)" + ) + cursor.execute( + "CREATE INDEX IF NOT EXISTS idx_index_composition_date ON index_composition(date)" + ) + cursor.execute( + "CREATE INDEX IF NOT EXISTS idx_features_ticker_date ON features(ticker, date)" + ) + + conn.commit() + + def _load_data(self): + """Load existing data from disk with memory management""" + try: + # Check memory before loading + if not self.memory_manager.ensure_memory(2 * 1024**3): # 2GB + logger.warning("Skipping data load due to memory constraints") + return + + self.price_data = self._load_pickle("price_data.pkl") + self.corporate_actions = self._load_pickle("corporate_actions.pkl") + self.sector_data = self._load_pickle("sector_data.pkl") + self.index_composition = self._load_pickle("index_composition.pkl") + + logger.info("Loaded existing data from disk") + self.memory_manager.log_memory_usage("[After Data Load]") + + except FileNotFoundError: + logger.info("No existing data found. Starting with empty databases.") + except Exception as e: + logger.error(f"Error loading data: {str(e)}", exc_info=True) + self.memory_manager.empty_cache() + + def _save_data(self): + """Save data to disk with memory management""" + try: + # Check memory before saving + if not self.memory_manager.ensure_memory(1 * 1024**3): # 1GB + logger.warning("Skipping data save due to memory constraints") + return + + self._save_pickle(self.price_data, "price_data.pkl") + self._save_pickle(self.corporate_actions, "corporate_actions.pkl") + self._save_pickle(self.sector_data, "sector_data.pkl") + self._save_pickle(self.index_composition, "index_composition.pkl") + + logger.info("Saved data to disk") + self.memory_manager.log_memory_usage("[After Data Save]") + + except Exception as e: + logger.error(f"Error saving data: {str(e)}", exc_info=True) + self.memory_manager.empty_cache() + + def _load_pickle(self, filename: str): + """Load data from pickle file with memory management""" + filepath = os.path.join(config.PROCESSED_DATA_DIR, filename) + if not os.path.exists(filepath): + raise FileNotFoundError(f"File {filepath} not found") + + with open(filepath, "rb") as f: + return pickle.load(f) + + def _save_pickle(self, data, filename: str): + """Save data to pickle file with memory management""" + filepath = os.path.join(config.PROCESSED_DATA_DIR, filename) + with open(filepath, "wb") as f: + pickle.dump(data, f) + + def _load_delisted_tickers(self): + """Load delisted tickers from external file""" + if os.path.exists(config.DELISTED_TICKERS_FILE): + try: + df = pd.read_csv(config.DELISTED_TICKERS_FILE) + self.delisted_tickers = set(df["Ticker"].tolist()) + logger.info(f"Loaded {len(self.delisted_tickers)} delisted tickers") + except Exception as e: + logger.error(f"Error loading delisted tickers: {str(e)}") + else: + logger.warning( + "Delisted tickers file not found. Only using active tickers." + ) + + def update_all_data(self): + """Update all data sources with memory management""" + # Get all tickers (active + delisted) + all_tickers = config.INITIAL_TICKERS + list(self.delisted_tickers) + + # Update price data + self.update_price_data(all_tickers) + + # Update corporate actions + self.update_corporate_actions(all_tickers) + + # Update sector data + self.update_sector_data(all_tickers) + + # Update index composition + self.update_index_composition() + + # Update alternative data + self.update_alternative_data(all_tickers) + + # Save updated data + self._save_data() + + def update_price_data(self, tickers: List[str]): + """Update price data for given tickers with memory management""" + logger.info(f"Updating price data for {len(tickers)} tickers") + + # Use yfinance for historical data + import yfinance as yf + + for ticker in tqdm(tickers, desc="Updating price data"): + try: + # Check memory before processing + if not self.memory_manager.ensure_memory(100 * 1024**2): # 100MB + logger.warning(f"Skipping {ticker} due to memory constraints") + continue + + # Determine start date + start_date = config.START_DATE + if ticker in self.price_data and not self.price_data[ticker].empty: + # If we already have data, start from the day after our last data point + start_date = ( + self.price_data[ticker].index[-1] + timedelta(days=1) + ).strftime("%Y-%m-%d") + + # Download new data + new_data = yf.download( + ticker, + start=start_date, + end=config.END_DATE, + progress=False, + auto_adjust=True, # Use adjusted prices + ) + + if not new_data.empty: + # If we already have data for this ticker, append the new data + if ticker in self.price_data and not self.price_data[ticker].empty: + # Combine existing and new data + combined = pd.concat([self.price_data[ticker], new_data]) + # Remove duplicates (keep the new data) + combined = combined[~combined.index.duplicated(keep="last")] + self.price_data[ticker] = combined.sort_index() + else: + self.price_data[ticker] = new_data.sort_index() + + # Store in database + self._store_price_data(ticker, new_data) + + # Memory management + self.memory_manager.auto_manage_memory(threshold=0.7) + + except Exception as e: + logger.warning(f"Error updating price data for {ticker}: {str(e)}") + self.memory_manager.empty_cache() + continue + + def _store_price_data(self, ticker: str, data: pd.DataFrame): + """Store price data in the database""" + if data.empty: + return + + # Convert to list of dictionaries for bulk insert + price_data = [] + for date, row in data.iterrows(): + price_data.append( + { + "ticker": ticker, + "date": date.strftime("%Y-%m-%d"), + "open": row["Open"], + "high": row["High"], + "low": row["Low"], + "close": row["Close"], + "adj_close": row["Adj Close"], + "volume": row["Volume"], + } + ) + + # Store in database + with sqlite3.connect(self.db_path) as conn: + cursor = conn.cursor() + cursor.executemany( + """ + INSERT OR REPLACE INTO price_data + (ticker, date, open, high, low, close, adj_close, volume) + VALUES (:ticker, :date, :open, :high, :low, :close, :adj_close, :volume) + """, + price_data, + ) + conn.commit() + + def update_corporate_actions(self, tickers: List[str]): + """Update corporate actions for given tickers with memory management""" + logger.info(f"Updating corporate actions for {len(tickers)} tickers") + + import yfinance as yf + + for ticker in tqdm(tickers, desc="Updating corporate actions"): + try: + # Check memory before processing + if not self.memory_manager.ensure_memory(50 * 1024**2): # 50MB + logger.warning(f"Skipping {ticker} due to memory constraints") + continue + + stock = yf.Ticker(ticker) + + # Initialize corporate actions dictionary if needed + if ticker not in self.corporate_actions: + self.corporate_actions[ticker] = { + "splits": {}, + "dividends": {}, + "mergers": [], + "spin-offs": [], + } + + # Get splits + splits = stock.splits + if not splits.empty: + for date, ratio in splits.items(): + date_str = date.strftime("%Y-%m-%d") + self.corporate_actions[ticker]["splits"][date_str] = float( + ratio + ) + + # Get dividends + dividends = stock.dividends + if not dividends.empty: + for date, amount in dividends.items(): + date_str = date.strftime("%Y-%m-%d") + self.corporate_actions[ticker]["dividends"][date_str] = float( + amount + ) + + # Store in database + self._store_corporate_actions(ticker) + + # Memory management + self.memory_manager.auto_manage_memory(threshold=0.7) + + except Exception as e: + logger.warning( + f"Error updating corporate actions for {ticker}: {str(e)}" + ) + self.memory_manager.empty_cache() + continue + + def _store_corporate_actions(self, ticker: str): + """Store corporate actions in the database""" + if ticker not in self.corporate_actions: + return + + actions = self.corporate_actions[ticker] + + # Store splits + for date, ratio in actions["splits"].items(): + with sqlite3.connect(self.db_path) as conn: + cursor = conn.cursor() + cursor.execute( + """ + INSERT OR REPLACE INTO corporate_actions + (ticker, date, action_type, value, details) + VALUES (?, ?, ?, ?, ?) + """, + (ticker, date, "split", ratio, f"Split ratio: {ratio}"), + ) + + # Store dividends + for date, amount in actions["dividends"].items(): + with sqlite3.connect(self.db_path) as conn: + cursor = conn.cursor() + cursor.execute( + """ + INSERT OR REPLACE INTO corporate_actions + (ticker, date, action_type, value, details) + VALUES (?, ?, ?, ?, ?) + """, + (ticker, date, "dividend", amount, f"Dividend amount: {amount}"), + ) + + def update_sector_data(self, tickers: List[str]): + """Update sector data for given tickers with memory management""" + logger.info(f"Updating sector data for {len(tickers)} tickers") + + import yfinance as yf + + for ticker in tqdm(tickers, desc="Updating sector data"): + try: + # Check memory before processing + if not self.memory_manager.ensure_memory(50 * 1024**2): # 50MB + logger.warning(f"Skipping {ticker} due to memory constraints") + continue + + stock = yf.Ticker(ticker) + info = stock.info + + if "sector" in info: + self.sector_data[ticker] = { + "sector": info["sector"], + "industry": info.get("industry", "Unknown"), + } + + # Store in database + with sqlite3.connect(self.db_path) as conn: + cursor = conn.cursor() + cursor.execute( + """ + INSERT OR REPLACE INTO sector_data + (ticker, date, sector, industry) + VALUES (?, ?, ?, ?) + """, + ( + ticker, + config.END_DATE, + info["sector"], + info.get("industry", "Unknown"), + ), + ) + + # Memory management + self.memory_manager.auto_manage_memory(threshold=0.7) + + except Exception as e: + logger.warning(f"Error updating sector data for {ticker}: {str(e)}") + self.memory_manager.empty_cache() + continue + + def update_index_composition(self): + """Update index composition with memory management""" + logger.info(f"Updating index composition for {config.INDEX_TICKER}") + + import yfinance as yf + + try: + # Check memory before processing + if not self.memory_manager.ensure_memory(100 * 1024**2): # 100MB + logger.warning( + "Skipping index composition update due to memory constraints" + ) + return + + # Get current constituents + index = yf.Ticker(config.INDEX_TICKER) + constituents = index.get_holders() + + if constituents is not None and not constituents.empty: + current_date = datetime.now().strftime("%Y-%m-%d") + + if config.INDEX_TICKER not in self.index_composition: + self.index_composition[config.INDEX_TICKER] = {} + + self.index_composition[config.INDEX_TICKER][current_date] = ( + constituents["Symbol"].tolist() + ) + + # Store in database + self._store_index_composition() + + # Memory management + self.memory_manager.auto_manage_memory(threshold=0.7) + + except Exception as e: + logger.warning(f"Error updating index composition: {str(e)}") + self.memory_manager.empty_cache() + + def _store_index_composition(self): + """Store index composition in the database""" + for index_ticker, composition in self.index_composition.items(): + for date, members in composition.items(): + for member in members: + with sqlite3.connect(self.db_path) as conn: + cursor = conn.cursor() + cursor.execute( + """ + INSERT OR REPLACE INTO index_composition + (index_ticker, date, member_ticker) + VALUES (?, ?, ?) + """, + (index_ticker, date, member), + ) + + def update_alternative_data(self, tickers: List[str]): + """ + Update alternative data sources for given tickers with memory management + + Parameters: + tickers: List of tickers to update + """ + # Calculate date range + end_date = datetime.now().strftime("%Y-%m-%d") + start_date = ( + datetime.now() + - timedelta( + days=max(config.NEWS_LOOKBACK_DAYS, config.SOCIAL_MEDIA_LOOKBACK_DAYS) + ) + ).strftime("%Y-%m-%d") + + logger.info(f"Updating alternative data from {start_date} to {end_date}") + + # Update news data + self.news_processor.fetch_news(tickers, start_date, end_date) + + # Update social media data + self.social_processor.fetch_twitter_data(tickers, start_date, end_date) + self.social_processor.fetch_reddit_data(tickers, start_date, end_date) + + def create_training_dataset(self) -> List: + """ + Create a training dataset with alternative data features and AMD optimizations + + Returns: + List of PyG Data objects + """ + from torch_geometric.data import Data + + logger.info("Creating training dataset with AMD optimizations") + dates = pd.date_range(config.START_DATE, config.TRAIN_END_DATE) + dataset = [] + + for date in tqdm(dates, desc="Creating dataset"): + # Check memory before processing date + if not self.memory_manager.ensure_memory(500 * 1024**2): # 500MB + logger.warning(f"Skipping {date.date()} due to memory constraints") + self.memory_manager.empty_cache() + continue + + # Get current universe of stocks + current_tickers = [] + for ticker in config.INITIAL_TICKERS + list(self.delisted_tickers): + if ticker in self.price_data and not self.price_data[ticker].empty: + if ( + date >= self.price_data[ticker].index[0] + and date <= self.price_data[ticker].index[-1] + ): + current_tickers.append(ticker) + + # Skip if no stocks available + if not current_tickers: + continue + + # Create node features + node_features = [] + corporate_action_flags = [] + + for ticker in current_tickers: + # Check memory before processing ticker + if not self.memory_manager.ensure_memory(10 * 1024**2): # 10MB + logger.warning(f"Skipping {ticker} due to memory constraints") + continue + + try: + # Get price data for lookback period + lookback_start = date - timedelta(days=config.LOOKBACK_WINDOW) + price_data = self.price_data[ticker].loc[lookback_start:date] + + if ( + len(price_data) < 5 + ): # Need at least 5 days for meaningful features + # Use default values + features = [ + 0, + 0, + 0, + 0, + 0, + ] # [return, volatility, momentum, volume, price] + else: + returns = price_data["Adj Close"].pct_change().dropna() + features = [ + returns.iloc[-1], # Last return + returns.std(), # Volatility + returns.mean(), # Momentum + np.log(price_data["Volume"].iloc[-1] + 1), # Log volume + price_data["Adj Close"].iloc[-1], # Last price + ] + + node_features.append(features) + + # Get corporate action flag + pit_data = self.get_point_in_time_data(ticker, date) + flag = 0 + if pit_data["upcoming_actions"]: + # Use the type of the soonest upcoming action + soonest = min( + pit_data["upcoming_actions"], + key=lambda x: x["days_until"], + ) + if soonest["type"] == "split": + flag = 1 + elif soonest["type"] == "dividend": + flag = 2 + + corporate_action_flags.append(flag) + + except Exception as e: + logger.warning( + f"Error processing {ticker} for {date.date()}: {str(e)}" + ) + continue + + # Skip if no features were created + if not node_features: + continue + + # Convert to tensors + x = torch.tensor(node_features, dtype=torch.float32) + + # Add corporate action flags as additional features + corporate_action_tensor = torch.tensor( + corporate_action_flags, dtype=torch.float32 + ).unsqueeze(1) + x = torch.cat([x, corporate_action_tensor], dim=1) + + # Create edges based on sector relationships + edge_index = [] + edge_weight = [] + + for i, ticker1 in enumerate(current_tickers): + for j, ticker2 in enumerate(current_tickers): + if i < j: + # Check memory before processing edge + if not self.memory_manager.ensure_memory(1 * 1024**2): # 1MB + logger.warning( + "Skipping edge creation due to memory constraints" + ) + continue + + try: + # Get sector data + pit1 = self.get_point_in_time_data(ticker1, date) + pit2 = self.get_point_in_time_data(ticker2, date) + + # Create edge if same sector + if ( + pit1["sector"] + and pit2["sector"] + and pit1["sector"] == pit2["sector"] + ): + # Calculate correlation as edge weight + lookback_start = date - timedelta( + days=config.LOOKBACK_WINDOW + ) + returns1 = ( + self.price_data[ticker1] + .loc[lookback_start:date]["Adj Close"] + .pct_change() + ) + returns2 = ( + self.price_data[ticker2] + .loc[lookback_start:date]["Adj Close"] + .pct_change() + ) + + if len(returns1) > 5 and len(returns2) > 5: + corr = returns1.corr(returns2) + if not np.isnan(corr): + edge_index.append([i, j]) + edge_weight.append(corr) + except Exception as e: + logger.warning( + f"Error creating edge between {ticker1} and {ticker2}: {str(e)}" + ) + continue + + # Convert to tensors + edge_index = ( + torch.tensor(edge_index, dtype=torch.long).t() + if edge_index + else torch.empty((2, 0), dtype=torch.long) + ) + edge_weight = ( + torch.tensor(edge_weight, dtype=torch.float32).unsqueeze(1) + if edge_weight + else torch.empty((0, 1), dtype=torch.float32) + ) + + # Create target (next day's return) + y = [] + for ticker in current_tickers: + next_date = date + timedelta(days=1) + if ( + ticker in self.price_data + and next_date in self.price_data[ticker].index + ): + ret = ( + self.price_data[ticker].loc[next_date]["Adj Close"] + / self.price_data[ticker].loc[date]["Adj Close"] + - 1 + ) + y.append(ret) + else: + y.append(0) # Default value + + y = torch.tensor(y, dtype=torch.float32).unsqueeze(1) + + # Create Data object + data = Data(x=x, edge_index=edge_index, edge_attr=edge_weight, y=y) + data.date = date + data.tickers = current_tickers + + dataset.append(data) + + # Memory management + self.memory_manager.auto_manage_memory(threshold=0.7) + + return dataset + + def create_intraday_dataset( + self, tickers: List[str], start_date: str, end_date: str + ) -> List: + """ + Create a dataset for intraday trading with AMD optimizations + + Parameters: + tickers: List of tickers + start_date: Start date (YYYY-MM-DD) + end_date: End date (YYYY-MM-DD) + + Returns: + List of PyG Data objects for intraday trading + """ + from torch_geometric.data import Data + + logger.info( + f"Creating intraday dataset from {start_date} to {end_date} with AMD optimizations" + ) + dataset = [] + + # Generate all trading days in the date range + date_range = pd.date_range(start_date, end_date) + trading_days = [ + date.strftime("%Y-%m-%d") for date in date_range if date.weekday() < 5 + ] # Weekdays only + + for date in trading_days: + # Check memory before processing date + if not self.memory_manager.ensure_memory(1 * 1024**3): # 1GB + logger.warning(f"Skipping {date} due to memory constraints") + self.memory_manager.empty_cache() + continue + + # Get all timestamps for this trading day + timestamps = generate_intraday_timestamps(date) + + for i in range(config.SEQUENCE_LENGTH, len(timestamps)): + current_timestamp = timestamps[i] + sequence_start = timestamps[i - config.SEQUENCE_LENGTH] + + # Get current universe of stocks + current_tickers = [] + for ticker in tickers: + if ticker in self.price_data and not self.price_data[ticker].empty: + if ( + date >= self.price_data[ticker].index[0] + and date <= self.price_data[ticker].index[-1] + ): + current_tickers.append(ticker) + + # Skip if no stocks available + if not current_tickers: + continue + + # Create node features for each stock in the sequence + sequence_features = [] + valid_tickers = [] + + for ticker in current_tickers: + # Check memory before processing ticker + if not self.memory_manager.ensure_memory(50 * 1024**2): # 50MB + logger.warning(f"Skipping {ticker} due to memory constraints") + continue + + try: + # Get price data for lookback period + lookback_start = datetime.strptime( + sequence_start, "%Y-%m-%d %H:%M:%S" + ) + lookback_end = datetime.strptime( + current_timestamp, "%Y-%m-%d %H:%M:%S" + ) + + # For intraday data, we need to get it from the live data service + # In a real implementation, this would come from the live data feed + # For this example, we'll simulate it with daily data + + # Get daily data for the date + daily_data = self.price_data[ticker].loc[date] + + if isinstance(daily_data, pd.Series): + # If only one day of data, create a sequence with the same values + features = [ + daily_data["Adj Close"] / daily_data["Open"] + - 1, # Return + 0.2, # Volatility (placeholder) + 0.0, # Momentum (placeholder) + np.log(daily_data["Volume"] + 1), # Log volume + daily_data["Adj Close"], # Price + ] + + # Repeat for the sequence + sequence_data = np.tile( + features, (config.SEQUENCE_LENGTH, 1) + ) + else: + # This shouldn't happen with daily data + continue + + sequence_features.append(sequence_data) + valid_tickers.append(ticker) + + except Exception as e: + logger.warning( + f"Error processing {ticker} for {current_timestamp}: {str(e)}" + ) + continue + + # Skip if no features were created + if not sequence_features: + continue + + # Convert to tensor (num_stocks, sequence_length, num_features) + x = torch.tensor(np.array(sequence_features), dtype=torch.float32) + + # Create edges based on sector relationships + edge_index = [] + edge_weight = [] + + for j, ticker1 in enumerate(valid_tickers): + for k, ticker2 in enumerate(valid_tickers): + if j < k: + # Check memory before processing edge + if not self.memory_manager.ensure_memory( + 1 * 1024**2 + ): # 1MB + logger.warning( + "Skipping edge creation due to memory constraints" + ) + continue + + try: + # Get sector relationship + pit1 = self.get_point_in_time_data( + ticker1, datetime.strptime(date, "%Y-%m-%d") + ) + pit2 = self.get_point_in_time_data( + ticker2, datetime.strptime(date, "%Y-%m-%d") + ) + + if ( + pit1["sector"] + and pit2["sector"] + and pit1["sector"] == pit2["sector"] + ): + # Calculate correlation of recent returns + # For intraday, we would use intraday returns + # For this example, we'll use daily returns + lookback_start = datetime.strptime( + date, "%Y-%m-%d" + ) - timedelta(days=config.LOOKBACK_WINDOW) + returns1 = ( + self.price_data[ticker1] + .loc[lookback_start:date]["Adj Close"] + .pct_change() + ) + returns2 = ( + self.price_data[ticker2] + .loc[lookback_start:date]["Adj Close"] + .pct_change() + ) + + if len(returns1) > 5 and len(returns2) > 5: + corr = returns1.corr(returns2) + if not np.isnan(corr): + edge_index.append([j, k]) + edge_weight.append(corr) + except Exception as e: + logger.warning( + f"Error creating edge between {ticker1} and {ticker2}: {str(e)}" + ) + continue + + # Convert to tensors + edge_index = ( + torch.tensor(edge_index, dtype=torch.long).t() + if edge_index + else torch.empty((2, 0), dtype=torch.long) + ) + edge_weight = ( + torch.tensor(edge_weight, dtype=torch.float32).unsqueeze(1) + if edge_weight + else torch.empty((0, 1), dtype=torch.float32) + ) + + # Create target (returns over prediction horizon) + y = [] + for ticker in valid_tickers: + # For intraday, we would predict the next few minutes + # For this example, we'll predict the next day's return + next_date = datetime.strptime(date, "%Y-%m-%d") + timedelta(days=1) + if ( + ticker in self.price_data + and next_date.strftime("%Y-%m-%d") + in self.price_data[ticker].index + ): + current_price = self.price_data[ticker].loc[date]["Adj Close"] + future_price = self.price_data[ticker].loc[ + next_date.strftime("%Y-%m-%d") + ]["Adj Close"] + ret = future_price / current_price - 1 + y.append(ret) + else: + y.append(0) # Default value + + y = torch.tensor(y, dtype=torch.float32).unsqueeze(1) + + # Create Data object + data = Data(x=x, edge_index=edge_index, edge_attr=edge_weight, y=y) + data.timestamp = current_timestamp + data.date = date + data.tickers = valid_tickers + data.sequence_start = sequence_start + data.sequence_end = current_timestamp + + dataset.append(data) + + # Memory management + self.memory_manager.auto_manage_memory(threshold=0.7) + + return dataset + + def get_point_in_time_data(self, ticker: str, date: datetime) -> Dict: + """ + Get point-in-time data for a stock at a specific date + + Parameters: + ticker: Stock ticker + date: Date as datetime object + + Returns: + Dictionary with point-in-time data + """ + date_str = date.strftime("%Y-%m-%d") + result = { + "ticker": ticker, + "date": date_str, + "price": None, + "sector": None, + "industry": None, + "in_index": False, + "index": None, + "upcoming_actions": [], + } + + # Get price data + if ticker in self.price_data and not self.price_data[ticker].empty: + # Find the most recent price before or on the date + price_data = self.price_data[ticker] + idx = price_data.index.get_indexer([date], method="ffill")[0] + if idx >= 0: + result["price"] = price_data.iloc[idx]["Adj Close"] + + # Get sector data + if ticker in self.sector_data: + result["sector"] = self.sector_data[ticker]["sector"] + result["industry"] = self.sector_data[ticker]["industry"] + + # Check if in index (simplified - would need historical index composition) + for index_ticker, composition in self.index_composition.items(): + # Find the most recent composition before the date + comp_dates = sorted(composition.keys()) + for comp_date in reversed(comp_dates): + if datetime.strptime(comp_date, "%Y-%m-%d") <= date: + if ticker in composition[comp_date]: + result["in_index"] = True + result["index"] = index_ticker + break + + # Get upcoming corporate actions + if ticker in self.corporate_actions: + actions = self.corporate_actions[ticker] + + # Check for upcoming splits (within 30 days) + for action_date, ratio in actions["splits"].items(): + action_date_dt = datetime.strptime(action_date, "%Y-%m-%d") + if date < action_date_dt <= date + timedelta(days=30): + result["upcoming_actions"].append( + { + "type": "split", + "date": action_date, + "ratio": ratio, + "days_until": (action_date_dt - date).days, + } + ) + + # Check for upcoming dividends (within 7 days) + for action_date, amount in actions["dividends"].items(): + action_date_dt = datetime.strptime(action_date, "%Y-%m-%d") + if date < action_date_dt <= date + timedelta(days=7): + result["upcoming_actions"].append( + { + "type": "dividend", + "date": action_date, + "amount": amount, + "days_until": (action_date_dt - date).days, + } + ) + + return result + + def get_latest_features( + self, tickers: List[str], timestamp: str + ) -> Dict[str, Dict]: + """ + Get the latest features for given tickers at a specific timestamp + + Parameters: + tickers: List of tickers + timestamp: Timestamp (YYYY-MM-DD HH:MM:SS) + + Returns: + Dictionary of {ticker: features} for the latest available data + """ + features = {} + + for ticker in tickers: + # Get price features + price_features = self._get_latest_price_features(ticker, timestamp) + + # Get alternative data features + news_features = self.news_processor.get_news_features( + ticker, timestamp.split(" ")[0] + ) + social_features = self.social_processor.get_all_social_features( + [ticker], timestamp.split(" ")[0] + ) + + # Combine features + if price_features: + combined_features = { + "ticker": ticker, + "timestamp": timestamp, + **price_features, + **news_features, + **( + social_features.iloc[0].to_dict() + if not social_features.empty + else {} + ), + } + features[ticker] = combined_features + + return features + + def _get_latest_price_features(self, ticker: str, timestamp: str) -> Dict: + """ + Get the latest price features for a ticker at a specific timestamp + + Parameters: + ticker: Stock ticker + timestamp: Timestamp (YYYY-MM-DD HH:MM:SS) + + Returns: + Dictionary of price features + """ + if ticker not in self.price_data or self.price_data[ticker].empty: + return {} + + # Get the date from the timestamp + date_str = timestamp.split(" ")[0] + date = datetime.strptime(date_str, "%Y-%m-%d") + + # Get price data for the date + if date_str not in self.price_data[ticker].index: + return {} + + price_data = self.price_data[ticker].loc[date_str] + + # Calculate features + features = { + "return": 0.0, # For intraday, this would be the return since open + "volatility": 0.2, # Placeholder - would calculate from intraday data + "momentum": 0.0, # Placeholder + "volume": np.log(price_data["Volume"] + 1), + "price": price_data["Adj Close"], + } + + return features + + def create_validation_dataset(self, start_date: str, end_date: str) -> List: + """ + Create a validation dataset (placeholder - reuses training logic) + """ + # For simplicity, reuse create_training_dataset logic with a different date range + # In a real implementation, this would be more sophisticated + return self.create_training_dataset() diff --git a/src/data/sentiment.py b/src/data/sentiment.py new file mode 100644 index 0000000..f9e04ec --- /dev/null +++ b/src/data/sentiment.py @@ -0,0 +1,31 @@ +""" +Sentiment analysis module for financial text. +""" + +import logging +from typing import Dict + +logger = logging.getLogger(__name__) + + +class SentimentAnalyzer: + """Analyze sentiment of financial text.""" + + def __init__(self): + # Placeholder: load a pre-trained sentiment model (e.g., FinBERT) + pass + + def analyze(self, text: str) -> Dict: + """Analyze sentiment of a text snippet.""" + # Placeholder: return neutral sentiment + return { + "label": "neutral", + "score": 0.0, + "positive": 0.33, + "negative": 0.33, + "neutral": 0.34, + } + + def analyze_batch(self, texts: list) -> list: + """Analyze sentiment of a batch of texts.""" + return [self.analyze(t) for t in texts] diff --git a/src/data/social_processor.py b/src/data/social_processor.py new file mode 100644 index 0000000..5fc6748 --- /dev/null +++ b/src/data/social_processor.py @@ -0,0 +1,59 @@ +""" +Social media data processor for fetching and analyzing social sentiment. +""" + +import logging +from typing import Dict, List + +import pandas as pd + +from config import config + +logger = logging.getLogger(__name__) + + +class SocialMediaProcessor: + """Process social media data for tickers.""" + + def __init__(self): + self.twitter_data = {} + self.reddit_data = {} + + def fetch_twitter_data(self, tickers: List[str], start_date: str, end_date: str): + """Fetch Twitter data for the given tickers and date range.""" + logger.info( + f"Fetching Twitter data for {len(tickers)} tickers from {start_date} to {end_date}" + ) + # Placeholder: integrate with Twitter API (e.g., Tweepy) + pass + + def fetch_reddit_data(self, tickers: List[str], start_date: str, end_date: str): + """Fetch Reddit data for the given tickers and date range.""" + logger.info( + f"Fetching Reddit data for {len(tickers)} tickers from {start_date} to {end_date}" + ) + # Placeholder: integrate with Reddit API (e.g., PRAW) + pass + + def get_all_social_features(self, tickers: List[str], date: str) -> pd.DataFrame: + """Return social media features for tickers on a specific date.""" + # Placeholder: return default features DataFrame + features = { + "ticker": tickers, + "twitter_sentiment": [0.0] * len(tickers), + "twitter_volume": [0] * len(tickers), + "reddit_sentiment": [0.0] * len(tickers), + "reddit_volume": [0] * len(tickers), + "social_momentum": [0.0] * len(tickers), + } + return pd.DataFrame(features) + + def get_social_features(self, ticker: str, date: str) -> Dict: + """Return social media features for a single ticker.""" + return { + "twitter_sentiment": 0.0, + "twitter_volume": 0, + "reddit_sentiment": 0.0, + "reddit_volume": 0, + "social_momentum": 0.0, + } diff --git a/src/evaluation/__init__.py b/src/evaluation/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/evaluation/backtester.py b/src/evaluation/backtester.py new file mode 100644 index 0000000..428dcfb --- /dev/null +++ b/src/evaluation/backtester.py @@ -0,0 +1,61 @@ +""" +Backtesting framework for the GNN trading strategy. +""" + +import logging +from typing import Dict, List, Tuple + +import pandas as pd + +from config import config +from src.models.trainer import GNNTrainer +from src.trading.paper_broker import PaperTradingBroker + +logger = logging.getLogger(__name__) + + +class GNNBacktester: + """Backtest the GNN model on historical data.""" + + def __init__(self, model, pipeline): + self.model = model + self.pipeline = pipeline + self.broker = PaperTradingBroker(initial_cash=config.INITIAL_CAPITAL) + + def run_backtest(self, dataset: List) -> Tuple[pd.Series, List]: + """Run a backtest on the given dataset.""" + logger.info(f"Starting backtest with {len(dataset)} samples") + + portfolio_values = [] + dates = [] + trade_log = [] + + for data in dataset: + # Get predictions + predictions = self.model(data) + + # Simulate trading based on predictions + for i, ticker in enumerate(data.tickers): + pred = predictions[i].item() + + if pred > 0.002: + order = { + "ticker": ticker, + "action": "buy", + "quantity": 100, + "price": 100, # placeholder + "timestamp": str(getattr(data, "date", "")), + "type": "market", + } + order_id = self.broker.submit_order(order) + if order_id: + trade_log.append({**order, "order_id": order_id}) + + # Record portfolio value + account = self.broker.get_account_summary() + portfolio_values.append(account["total_value"]) + dates.append(getattr(data, "date", None)) + + portfolio_series = pd.Series(portfolio_values, index=dates) + logger.info("Backtest completed") + return portfolio_series, trade_log diff --git a/src/evaluation/intraday_backtester.py b/src/evaluation/intraday_backtester.py new file mode 100644 index 0000000..4c52cf6 --- /dev/null +++ b/src/evaluation/intraday_backtester.py @@ -0,0 +1,57 @@ +""" +Intraday backtesting framework. +""" + +import logging +from typing import Dict, List, Tuple + +import pandas as pd + +from config import config +from src.trading.paper_broker import PaperTradingBroker + +logger = logging.getLogger(__name__) + + +class IntradayBacktester: + """Backtest intraday trading strategies.""" + + def __init__(self, model, pipeline): + self.model = model + self.pipeline = pipeline + self.broker = PaperTradingBroker(initial_cash=config.INITIAL_CAPITAL) + + def run_backtest(self, dataset: List) -> Tuple[pd.Series, List]: + """Run an intraday backtest.""" + logger.info(f"Starting intraday backtest with {len(dataset)} samples") + + portfolio_values = [] + timestamps = [] + trade_log = [] + + for data in dataset: + predictions = self.model(data) + + for i, ticker in enumerate(getattr(data, "tickers", [])): + pred = predictions[i].item() + + if pred > 0.002: + order = { + "ticker": ticker, + "action": "buy", + "quantity": 100, + "price": 100, + "timestamp": getattr(data, "timestamp", ""), + "type": "market", + } + order_id = self.broker.submit_order(order) + if order_id: + trade_log.append({**order, "order_id": order_id}) + + account = self.broker.get_account_summary() + portfolio_values.append(account["total_value"]) + timestamps.append(getattr(data, "timestamp", None)) + + portfolio_series = pd.Series(portfolio_values, index=timestamps) + logger.info("Intraday backtest completed") + return portfolio_series, trade_log diff --git a/src/evaluation/metrics.py b/src/evaluation/metrics.py new file mode 100644 index 0000000..0dcce1b --- /dev/null +++ b/src/evaluation/metrics.py @@ -0,0 +1,128 @@ +""" +Performance metrics for trading strategy evaluation. +""" + +import logging +from typing import Dict + +import numpy as np +import pandas as pd + +logger = logging.getLogger(__name__) + + +def calculate_performance_metrics( + portfolio_returns: pd.Series, benchmark_returns: pd.Series +) -> Dict: + """Calculate portfolio performance metrics.""" + metrics = {} + + # Total return + metrics["total_return"] = (1 + portfolio_returns).prod() - 1 + + # Annualized return + metrics["annualized_return"] = (1 + metrics["total_return"]) ** ( + 252 / len(portfolio_returns) + ) - 1 + + # Volatility + metrics["volatility"] = portfolio_returns.std() * np.sqrt(252) + + # Sharpe ratio (assuming risk-free rate of 0) + metrics["sharpe_ratio"] = ( + metrics["annualized_return"] / metrics["volatility"] + if metrics["volatility"] > 0 + else 0 + ) + + # Sortino ratio + downside_returns = portfolio_returns[portfolio_returns < 0] + downside_std = ( + downside_returns.std() * np.sqrt(252) if len(downside_returns) > 0 else 1e-6 + ) + metrics["sortino_ratio"] = metrics["annualized_return"] / downside_std + + # Maximum drawdown + cumulative = (1 + portfolio_returns).cumprod() + running_max = cumulative.expanding().max() + drawdown = (cumulative - running_max) / running_max + metrics["max_drawdown"] = drawdown.min() + + # Calmar ratio + metrics["calmar_ratio"] = ( + metrics["annualized_return"] / abs(metrics["max_drawdown"]) + if metrics["max_drawdown"] != 0 + else 0 + ) + + # Win rate + metrics["win_rate"] = (portfolio_returns > 0).mean() + + # Profit factor + gross_profit = portfolio_returns[portfolio_returns > 0].sum() + gross_loss = abs(portfolio_returns[portfolio_returns < 0].sum()) + metrics["profit_factor"] = ( + gross_profit / gross_loss if gross_loss > 0 else float("inf") + ) + + # Beta + covariance = portfolio_returns.cov(benchmark_returns) + benchmark_variance = benchmark_returns.var() + metrics["beta"] = covariance / benchmark_variance if benchmark_variance > 0 else 0 + + # Alpha + metrics["alpha"] = ( + metrics["annualized_return"] - metrics["beta"] * benchmark_returns.mean() * 252 + ) + + return metrics + + +def compare_to_benchmark( + portfolio_values: pd.Series, benchmark_values: pd.Series +) -> Dict: + """Compare portfolio performance to benchmark.""" + comparison = {} + + portfolio_returns = portfolio_values.pct_change().dropna() + benchmark_returns = benchmark_values.pct_change().dropna() + + # Total return comparison + comparison["portfolio_total_return"] = ( + portfolio_values.iloc[-1] / portfolio_values.iloc[0] + ) - 1 + comparison["benchmark_total_return"] = ( + benchmark_values.iloc[-1] / benchmark_values.iloc[0] + ) - 1 + comparison["excess_return"] = ( + comparison["portfolio_total_return"] - comparison["benchmark_total_return"] + ) + + # Tracking error + comparison["tracking_error"] = ( + portfolio_returns - benchmark_returns + ).std() * np.sqrt(252) + + # Information ratio + active_returns = portfolio_returns - benchmark_returns + comparison["information_ratio"] = ( + active_returns.mean() * 252 / (active_returns.std() * np.sqrt(252)) + if active_returns.std() > 0 + else 0 + ) + + # Up/down capture + up_market = benchmark_returns > 0 + down_market = benchmark_returns < 0 + comparison["up_capture"] = ( + (portfolio_returns[up_market].mean() / benchmark_returns[up_market].mean()) + if up_market.sum() > 0 + else 0 + ) + comparison["down_capture"] = ( + (portfolio_returns[down_market].mean() / benchmark_returns[down_market].mean()) + if down_market.sum() > 0 + else 0 + ) + + return comparison diff --git a/src/models/__init__.py b/src/models/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/models/gnn_model.py b/src/models/gnn_model.py new file mode 100644 index 0000000..c3a8660 --- /dev/null +++ b/src/models/gnn_model.py @@ -0,0 +1,383 @@ +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F +from torch_geometric.nn import GATConv + +from config import config +from src.amd.optimizations import AMDGATConv, AMDSparseAttention + + +class TemporalAttention(nn.Module): + """ + Temporal attention optimized for AMD GPUs + """ + + def __init__(self, feature_dim: int, num_heads: int = 8, dropout: float = 0.1): + super().__init__() + self.feature_dim = feature_dim + self.num_heads = num_heads + self.head_dim = feature_dim // num_heads + self.dropout = dropout + + # Use AMD-optimized sparse attention if available + try: + self.attention = AMDSparseAttention(feature_dim, num_heads, dropout) + except Exception: + # Fallback to standard multi-head attention + self.query = nn.Linear(feature_dim, feature_dim) + self.key = nn.Linear(feature_dim, feature_dim) + self.value = nn.Linear(feature_dim, feature_dim) + self.out = nn.Linear(feature_dim, feature_dim) + + self.layer_norm = nn.LayerNorm(feature_dim) + self.dropout_layer = nn.Dropout(dropout) + + # Initialize weights + self._init_weights() + + def _init_weights(self): + """Initialize weights with Xavier initialization""" + if not hasattr(self, "attention"): + nn.init.xavier_uniform_(self.query.weight) + nn.init.xavier_uniform_(self.key.weight) + nn.init.xavier_uniform_(self.value.weight) + nn.init.xavier_uniform_(self.out.weight) + + nn.init.zeros_(self.query.bias) + nn.init.zeros_(self.key.bias) + nn.init.zeros_(self.value.bias) + nn.init.zeros_(self.out.bias) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """ + Apply temporal attention to input sequence + + Parameters: + x: Input tensor of shape (batch_size, sequence_length, feature_dim) + + Returns: + Tensor of shape (batch_size, feature_dim) with temporal attention applied + """ + # Layer normalization + x_norm = self.layer_norm(x) + + if hasattr(self, "attention"): + # Use AMD-optimized sparse attention + attended = self.attention(x_norm) + else: + # Standard multi-head attention + batch_size, seq_len, _ = x_norm.size() + + # Project to query, key, value + q = self.query(x_norm) + k = self.key(x_norm) + v = self.value(x_norm) + + # Reshape for multi-head attention + q = q.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose( + 1, 2 + ) + k = k.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose( + 1, 2 + ) + v = v.view(batch_size, seq_len, self.num_heads, self.head_dim).transpose( + 1, 2 + ) + + # Calculate attention scores + attn_scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt( + self.head_dim + ) + + # Apply softmax + attn_weights = F.softmax(attn_scores, dim=-1) + attn_weights = self.dropout_layer(attn_weights) + + # Apply attention to values + attended = torch.matmul(attn_weights, v) + + # Concatenate heads + attended = ( + attended.transpose(1, 2) + .contiguous() + .view(batch_size, seq_len, self.feature_dim) + ) + + # Final projection + attended = self.out(attended) + + # Aggregate across time with residual connection + output = attended.mean(dim=1) + x.mean(dim=1) + + return output + + +class CorporateActionAwareGNN(nn.Module): + """ + GNN model with corporate action awareness, optimized for AMD Radeon R9700 AI Pro + """ + + def __init__(self, num_node_features: int): + super(CorporateActionAwareGNN, self).__init__() + + # Calculate feature dimensions + num_price_features = 5 # returns, volatility, momentum, volume, price + num_news_features = len(config.NEWS_FEATURES) + num_social_features = len(config.SOCIAL_FEATURES) + + # Feature processing modules with AMD optimizations + self.price_processor = nn.Sequential( + nn.Linear(num_price_features, config.HIDDEN_CHANNELS), + nn.SiLU(), # Swish activation often works better on AMD GPUs + nn.Linear(config.HIDDEN_CHANNELS, config.HIDDEN_CHANNELS), + nn.LayerNorm(config.HIDDEN_CHANNELS), + ) + + self.news_processor = nn.Sequential( + nn.Linear(num_news_features, config.HIDDEN_CHANNELS), + nn.SiLU(), + nn.Linear(config.HIDDEN_CHANNELS, config.HIDDEN_CHANNELS), + nn.LayerNorm(config.HIDDEN_CHANNELS), + ) + + self.social_processor = nn.Sequential( + nn.Linear(num_social_features, config.HIDDEN_CHANNELS), + nn.SiLU(), + nn.Linear(config.HIDDEN_CHANNELS, config.HIDDEN_CHANNELS), + nn.LayerNorm(config.HIDDEN_CHANNELS), + ) + + self.corporate_action_mlp = nn.Sequential( + nn.Linear(1, config.HIDDEN_CHANNELS), + nn.SiLU(), + nn.Linear(config.HIDDEN_CHANNELS, config.HIDDEN_CHANNELS), + nn.LayerNorm(config.HIDDEN_CHANNELS), + ) + + # Temporal attention for sequence processing + self.temporal_attention = TemporalAttention(config.HIDDEN_CHANNELS, num_heads=8) + + # Graph attention layers with AMD optimizations + try: + self.conv1 = AMDGATConv( + config.HIDDEN_CHANNELS * 4, # Combined features from all processors + config.HIDDEN_CHANNELS, + heads=config.NUM_HEADS, + concat=True, + dropout=config.DROPOUT, + add_self_loops=True, + ) + except Exception: + self.conv1 = GATConv( + config.HIDDEN_CHANNELS * 4, + config.HIDDEN_CHANNELS, + heads=config.NUM_HEADS, + concat=True, + dropout=config.DROPOUT, + add_self_loops=True, + ) + + try: + self.conv2 = AMDGATConv( + config.HIDDEN_CHANNELS * config.NUM_HEADS, + config.HIDDEN_CHANNELS, + heads=1, + concat=False, + dropout=config.DROPOUT, + add_self_loops=True, + ) + except Exception: + self.conv2 = GATConv( + config.HIDDEN_CHANNELS * config.NUM_HEADS, + config.HIDDEN_CHANNELS, + heads=1, + concat=False, + dropout=config.DROPOUT, + add_self_loops=True, + ) + + # Attention mechanism for combining alternative data + self.alternative_data_attention = nn.Sequential( + nn.Linear(config.HIDDEN_CHANNELS * 3, config.HIDDEN_CHANNELS), + nn.SiLU(), + nn.Linear( + config.HIDDEN_CHANNELS, 3 + ), # 3 attention weights (price, news, social) + nn.Softmax(dim=1), + ) + + # LSTM for temporal dependencies with AMD optimizations + self.lstm = nn.LSTM( + input_size=config.HIDDEN_CHANNELS, + hidden_size=config.HIDDEN_CHANNELS, + num_layers=2, + batch_first=True, + dropout=config.DROPOUT if config.NUM_HEADS > 1 else 0, + bidirectional=False, + ) + + # Initialize LSTM weights + self._init_lstm_weights() + + # Final prediction layer + self.linear = nn.Linear( + config.HIDDEN_CHANNELS * 2, 1 + ) # *2 for concatenating GNN output and attention features + + # Initialize weights + self._init_weights() + + def _init_weights(self): + """Initialize weights with AMD-friendly initialization""" + # Initialize linear layer weights + nn.init.xavier_uniform_(self.linear.weight) + nn.init.zeros_(self.linear.bias) + + # Initialize attention layers + for layer in self.alternative_data_attention: + if isinstance(layer, nn.Linear): + nn.init.xavier_uniform_(layer.weight) + nn.init.zeros_(layer.bias) + + def _init_lstm_weights(self): + """Initialize LSTM weights with orthogonal initialization for better convergence""" + for name, param in self.lstm.named_parameters(): + if "weight_ih" in name: + nn.init.orthogonal_(param) + elif "weight_hh" in name: + nn.init.orthogonal_(param) + elif "bias" in name: + nn.init.zeros_(param) + # Set forget gate bias to 1 for better gradient flow + n = param.size(0) + param.data[n // 4 : n // 2].fill_(1) + + def forward(self, data): + x, edge_index, edge_attr = data.x, data.edge_index, data.edge_attr + + # Split features into different types + num_price_features = 5 + num_news_features = len(config.NEWS_FEATURES) + num_social_features = len(config.SOCIAL_FEATURES) + + # Price features (first 5) + price_features = x[:, :, :num_price_features] + price_processed = self.price_processor(price_features) + + # News features (next num_news_features) + news_start = num_price_features + news_end = news_start + num_news_features + news_features = x[:, :, news_start:news_end] + news_processed = self.news_processor(news_features) + + # Social features (next num_social_features) + social_start = news_end + social_end = social_start + num_social_features + social_features = x[:, :, social_start:social_end] + social_processed = self.social_processor(social_features) + + # Corporate action flag (last feature) + corporate_action_flags = x[:, :, -1:] + corporate_action_features = self.corporate_action_mlp(corporate_action_flags) + + # Apply temporal attention to each feature type + price_attended = self.temporal_attention(price_processed) + news_attended = self.temporal_attention(news_processed) + social_attended = self.temporal_attention(social_processed) + corporate_attended = self.temporal_attention(corporate_action_features) + + # Calculate attention weights for alternative data + alternative_features = torch.cat( + [price_attended, news_attended, social_attended], dim=1 + ) + + attention_weights = self.alternative_data_attention(alternative_features) + + # Apply attention weights + weighted_price = price_attended * attention_weights[:, 0].unsqueeze(1) + weighted_news = news_attended * attention_weights[:, 1].unsqueeze(1) + weighted_social = social_attended * attention_weights[:, 2].unsqueeze(1) + + # Combine features + combined_features = torch.cat( + [weighted_price, weighted_news, weighted_social, corporate_attended], dim=1 + ) + + # Process through GNN with gradient checkpointing for memory efficiency + x = self._gnn_forward(combined_features, edge_index, edge_attr) + + # Process through LSTM for temporal dependencies + lstm_input = x.unsqueeze(1) # (num_stocks, 1, hidden_size) + lstm_out, _ = self.lstm(lstm_input) + lstm_out = lstm_out.squeeze(1) + + # Combine GNN output with attention features + attention_features = torch.cat( + [ + attention_weights[:, 0].unsqueeze(1), + attention_weights[:, 1].unsqueeze(1), + attention_weights[:, 2].unsqueeze(1), + ], + dim=1, + ) + + x = torch.cat([lstm_out, attention_features], dim=1) + + # Final prediction + return self.linear(x) + + def _gnn_forward(self, x, edge_index, edge_attr): + """Forward pass through GNN with gradient checkpointing for memory efficiency""" + # Gradient checkpointing for memory efficiency on AMD GPUs + x = torch.utils.checkpoint.checkpoint( + self.conv1, x, edge_index, edge_attr, preserve_rng_state=False + ) + x = F.silu(x) # Swish activation often works better than ReLU on AMD GPUs + x = F.dropout(x, p=config.DROPOUT, training=self.training) + + x = torch.utils.checkpoint.checkpoint( + self.conv2, x, edge_index, edge_attr, preserve_rng_state=False + ) + + return x + + def get_attention_weights(self, data): + """Get attention weights for interpretability""" + x, edge_index, edge_attr = data.x, data.edge_index, data.edge_attr + + # Process through the model up to attention + num_price_features = 5 + num_news_features = len(config.NEWS_FEATURES) + num_social_features = len(config.SOCIAL_FEATURES) + + price_features = x[:, :, :num_price_features] + news_features = x[ + :, :, num_price_features : num_price_features + num_news_features + ] + social_features = x[ + :, + :, + num_price_features + num_news_features : num_price_features + + num_news_features + + num_social_features, + ] + corporate_action_flags = x[:, :, -1:] + + price_processed = self.price_processor(price_features) + news_processed = self.news_processor(news_features) + social_processed = self.social_processor(social_features) + corporate_action_features = self.corporate_action_mlp(corporate_action_flags) + + price_attended = self.temporal_attention(price_processed) + news_attended = self.temporal_attention(news_processed) + social_attended = self.temporal_attention(social_processed) + + alternative_features = torch.cat( + [price_attended, news_attended, social_attended], dim=1 + ) + + attention_weights = self.alternative_data_attention(alternative_features) + + return attention_weights diff --git a/src/models/intraday_gnn.py b/src/models/intraday_gnn.py new file mode 100644 index 0000000..5791da6 --- /dev/null +++ b/src/models/intraday_gnn.py @@ -0,0 +1,220 @@ +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F +from torch_geometric.nn import GATConv + +from config import config +from src.amd.optimizations import AMDGATConv, AMDSparseAttention +from src.models.gnn_model import TemporalAttention + + +class IntradayGNN(nn.Module): + """ + Intraday GNN model optimized for AMD Radeon R9700 AI Pro + """ + + def __init__( + self, num_features: int, sequence_length: int = config.SEQUENCE_LENGTH + ): + super(IntradayGNN, self).__init__() + + self.sequence_length = sequence_length + self.num_features = num_features + + # Temporal attention for sequence processing + self.temporal_attention = TemporalAttention(num_features, num_heads=8) + + # Feature processing modules with AMD optimizations + self.feature_processor = nn.Sequential( + nn.Linear(num_features, config.HIDDEN_CHANNELS), + nn.SiLU(), # Swish activation often works better on AMD GPUs + nn.Linear(config.HIDDEN_CHANNELS, config.HIDDEN_CHANNELS), + nn.LayerNorm(config.HIDDEN_CHANNELS), + ) + + # Graph attention layers with AMD optimizations + try: + self.conv1 = AMDGATConv( + config.HIDDEN_CHANNELS, + config.HIDDEN_CHANNELS, + heads=config.NUM_HEADS, + concat=True, + dropout=config.DROPOUT, + add_self_loops=True, + ) + except Exception: + self.conv1 = GATConv( + config.HIDDEN_CHANNELS, + config.HIDDEN_CHANNELS, + heads=config.NUM_HEADS, + concat=True, + dropout=config.DROPOUT, + add_self_loops=True, + ) + + try: + self.conv2 = AMDGATConv( + config.HIDDEN_CHANNELS * config.NUM_HEADS, + config.HIDDEN_CHANNELS, + heads=1, + concat=False, + dropout=config.DROPOUT, + add_self_loops=True, + ) + except Exception: + self.conv2 = GATConv( + config.HIDDEN_CHANNELS * config.NUM_HEADS, + config.HIDDEN_CHANNELS, + heads=1, + concat=False, + dropout=config.DROPOUT, + add_self_loops=True, + ) + + # LSTM for temporal dependencies with AMD optimizations + self.lstm = nn.LSTM( + input_size=config.HIDDEN_CHANNELS, + hidden_size=config.HIDDEN_CHANNELS, + num_layers=2, + batch_first=True, + dropout=config.DROPOUT if config.NUM_HEADS > 1 else 0, + bidirectional=False, + ) + + # Initialize LSTM weights + self._init_lstm_weights() + + # Attention mechanism for final prediction + self.attention = nn.Sequential( + nn.Linear(config.HIDDEN_CHANNELS, config.HIDDEN_CHANNELS), + nn.SiLU(), + nn.Linear(config.HIDDEN_CHANNELS, 1), + nn.Softmax(dim=1), + ) + + # Final prediction layer + self.linear = nn.Linear(config.HIDDEN_CHANNELS, 1) + + # State for online learning + self.hidden_state = None + + # Initialize weights + self._init_weights() + + def _init_weights(self): + """Initialize weights with AMD-friendly initialization""" + # Initialize linear layer weights + nn.init.xavier_uniform_(self.linear.weight) + nn.init.zeros_(self.linear.bias) + + # Initialize attention layers + for layer in self.attention: + if isinstance(layer, nn.Linear): + nn.init.xavier_uniform_(layer.weight) + nn.init.zeros_(layer.bias) + + def _init_lstm_weights(self): + """Initialize LSTM weights with orthogonal initialization for better convergence""" + for name, param in self.lstm.named_parameters(): + if "weight_ih" in name: + nn.init.orthogonal_(param) + elif "weight_hh" in name: + nn.init.orthogonal_(param) + elif "bias" in name: + nn.init.zeros_(param) + # Set forget gate bias to 1 for better gradient flow + n = param.size(0) + param.data[n // 4 : n // 2].fill_(1) + + def forward(self, data): + x, edge_index, edge_attr = data.x, data.edge_index, data.edge_attr + + # x shape: (num_stocks, sequence_length, num_features) + batch_size, seq_len, num_features = x.size() + + # Apply temporal attention to each stock's sequence + temporal_features = [] + for i in range(batch_size): + stock_sequence = x[i].unsqueeze(0) # (1, sequence_length, num_features) + temporal_feature = self.temporal_attention(stock_sequence) + temporal_features.append(temporal_feature) + + # Stack temporal features + temporal_features = torch.cat( + temporal_features, dim=0 + ) # (num_stocks, feature_dim) + + # Process features + processed_features = self.feature_processor(temporal_features) + + # Process through GNN with gradient checkpointing for memory efficiency + x = self._gnn_forward(processed_features, edge_index, edge_attr) + + # Process through LSTM for temporal dependencies + # Reshape for LSTM: (num_stocks, 1, hidden_size) + lstm_input = x.unsqueeze(1) + + # If we have a hidden state from previous prediction, use it + if self.hidden_state is not None and config.STATEFUL_PREDICTION: + lstm_out, self.hidden_state = self.lstm(lstm_input, self.hidden_state) + else: + lstm_out, self.hidden_state = self.lstm(lstm_input) + + # Remove sequence dimension + lstm_out = lstm_out.squeeze(1) + + # Apply attention to LSTM outputs + attention_weights = self.attention(lstm_out) + attended = (lstm_out * attention_weights).sum(dim=1, keepdim=True) + + # Final prediction + return self.linear(attended) + + def _gnn_forward(self, x, edge_index, edge_attr): + """Forward pass through GNN with gradient checkpointing for memory efficiency""" + # Gradient checkpointing for memory efficiency on AMD GPUs + x = torch.utils.checkpoint.checkpoint( + self.conv1, x, edge_index, edge_attr, preserve_rng_state=False + ) + x = F.silu(x) # Swish activation often works better than ReLU on AMD GPUs + x = F.dropout(x, p=config.DROPOUT, training=self.training) + + x = torch.utils.checkpoint.checkpoint( + self.conv2, x, edge_index, edge_attr, preserve_rng_state=False + ) + + return x + + def reset_state(self): + """Reset the hidden state of the LSTM""" + self.hidden_state = None + + def get_attention_weights(self, data): + """Get attention weights for interpretability""" + x, edge_index, edge_attr = data.x, data.edge_index, data.edge_attr + + # Process through the model up to attention + batch_size, seq_len, num_features = x.size() + temporal_features = [] + + for i in range(batch_size): + stock_sequence = x[i].unsqueeze(0) + temporal_feature = self.temporal_attention(stock_sequence) + temporal_features.append(temporal_feature) + + temporal_features = torch.cat(temporal_features, dim=0) + processed_features = self.feature_processor(temporal_features) + + x = self.conv1(processed_features, edge_index, edge_attr) + x = F.silu(x) + x = self.conv2(x, edge_index, edge_attr) + + lstm_input = x.unsqueeze(1) + lstm_out, _ = self.lstm(lstm_input) + lstm_out = lstm_out.squeeze(1) + + attention_weights = self.attention(lstm_out) + + return attention_weights diff --git a/src/models/trainer.py b/src/models/trainer.py new file mode 100644 index 0000000..c72544b --- /dev/null +++ b/src/models/trainer.py @@ -0,0 +1,356 @@ +import logging +import os +import time +from typing import Dict, List, Tuple + +import numpy as np +import torch +import torch.nn as nn +from torch.cuda.amp import GradScaler, autocast +from torch_geometric.loader import DataLoader +from tqdm import tqdm + +from config import config +from src.amd.optimizations import AMDOptimizer +from src.utils.memory_manager import MemoryManager + +logger = logging.getLogger(__name__) + + +class GNNTrainer: + def __init__(self, model: nn.Module): + self.model = model + self.device = torch.device(config.DEVICE) + self.memory_manager = MemoryManager() + self.amd_optimizer = AMDOptimizer() + + # Optimize model for AMD GPU + self.model = self.amd_optimizer.optimize_model(self.model) + + # Set up optimizer with weight decay + self.optimizer = torch.optim.AdamW( + self.model.parameters(), lr=config.LEARNING_RATE, weight_decay=1e-4 + ) + + # Set up learning rate scheduler + self.scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau( + self.optimizer, mode="min", factor=0.5, patience=5, verbose=True + ) + + self.criterion = nn.MSELoss() + + # For online learning + if config.ONLINE_LEARNING: + self.online_optimizer = torch.optim.AdamW( + self.model.parameters(), lr=config.ONLINE_LEARNING_RATE + ) + + # For mixed precision training + self.scaler = GradScaler(enabled=config.MIXED_PRECISION) + + # Training statistics + self.train_losses = [] + self.val_losses = [] + self.best_val_loss = float("inf") + + # Log memory info + logger.info(self.memory_manager.get_memory_stats()) + + def train( + self, train_dataset: List, val_dataset: List + ) -> Tuple[List[float], List[float]]: + """ + Train the GNN model with AMD optimizations + + Parameters: + train_dataset: Training dataset + val_dataset: Validation dataset + + Returns: + Tuple of (train_losses, val_losses) + """ + # Create data loaders with AMD optimizations + train_loader = DataLoader( + train_dataset, + batch_size=config.BATCH_SIZE, + shuffle=True, + num_workers=config.NUM_WORKERS, + pin_memory=config.PIN_MEMORY, + prefetch_factor=config.PREFETCH_FACTOR, + ) + + val_loader = DataLoader( + val_dataset, + batch_size=config.BATCH_SIZE, + shuffle=False, + num_workers=config.NUM_WORKERS, + pin_memory=config.PIN_MEMORY, + prefetch_factor=config.PREFETCH_FACTOR, + ) + + for epoch in range(config.EPOCHS): + # Training + self.model.train() + epoch_train_loss = 0.0 + start_time = time.time() + + for batch in tqdm( + train_loader, desc=f"Epoch {epoch + 1}/{config.EPOCHS} - Training" + ): + # Check memory before processing batch + if not self.memory_manager.ensure_memory(2 * 1024**3): # 2GB + logger.warning("Skipping batch due to memory constraints") + continue + + try: + batch = batch.to(self.device, non_blocking=config.PIN_MEMORY) + self.optimizer.zero_grad( + set_to_none=True + ) # More efficient for AMD GPUs + + # Mixed precision training + with autocast( + enabled=config.MIXED_PRECISION, + dtype=self.amd_optimizer.get_precision_dtype(), + ): + out = self.model(batch) + loss = self.criterion(out, batch.y) + + # Scale loss and backpropagate + self.scaler.scale(loss).backward() + + # Gradient clipping for stability + self.scaler.unscale_(self.optimizer) + torch.nn.utils.clip_grad_norm_( + self.model.parameters(), max_norm=1.0 + ) + + # Update weights + self.scaler.step(self.optimizer) + self.scaler.update() + + epoch_train_loss += loss.item() + + # Memory management + self.memory_manager.auto_manage_memory(threshold=0.8) + + except Exception as e: + logger.error(f"Error processing batch: {str(e)}", exc_info=True) + self.memory_manager.empty_cache() + continue + + epoch_train_loss /= len(train_loader) + self.train_losses.append(epoch_train_loss) + + # Validation + self.model.eval() + epoch_val_loss = 0.0 + + with torch.no_grad(): + for batch in val_loader: + # Check memory before processing batch + if not self.memory_manager.ensure_memory(1 * 1024**3): # 1GB + logger.warning( + "Skipping validation batch due to memory constraints" + ) + continue + + try: + batch = batch.to(self.device, non_blocking=config.PIN_MEMORY) + + with autocast( + enabled=config.MIXED_PRECISION, + dtype=self.amd_optimizer.get_precision_dtype(), + ): + out = self.model(batch) + loss = self.criterion(out, batch.y) + + epoch_val_loss += loss.item() + + # Memory management + self.memory_manager.auto_manage_memory(threshold=0.8) + + except Exception as e: + logger.error( + f"Error processing validation batch: {str(e)}", + exc_info=True, + ) + self.memory_manager.empty_cache() + continue + + epoch_val_loss /= len(val_loader) + self.val_losses.append(epoch_val_loss) + + # Update learning rate scheduler + self.scheduler.step(epoch_val_loss) + + # Log training information + epoch_time = time.time() - start_time + memory_info = self.memory_manager.check_memory() + + logger.info( + f"Epoch {epoch + 1}/{config.EPOCHS} - " + f"Train Loss: {epoch_train_loss:.6f}, " + f"Val Loss: {epoch_val_loss:.6f}, " + f"LR: {self.optimizer.param_groups[0]['lr']:.2e}, " + f"Time: {epoch_time:.2f}s, " + f"Memory: {memory_info['allocated'] / 1024**3:.2f}GB/{memory_info['limit'] / 1024**3:.2f}GB" + ) + + # Save best model + if epoch_val_loss < self.best_val_loss: + self.best_val_loss = epoch_val_loss + self.save_model() + logger.info("Saved best model") + + return self.train_losses, self.val_losses + + def online_update(self, data): + """ + Perform online learning update with new data using AMD optimizations + + Parameters: + data: New data for online learning + """ + if not config.ONLINE_LEARNING: + return + + # Check memory before online learning + if not self.memory_manager.ensure_memory(1 * 1024**3): # 1GB + logger.warning("Skipping online learning due to memory constraints") + return + + self.model.train() + self.online_optimizer.zero_grad(set_to_none=True) + + try: + data = data.to(self.device, non_blocking=config.PIN_MEMORY) + + with autocast( + enabled=config.MIXED_PRECISION, + dtype=self.amd_optimizer.get_precision_dtype(), + ): + out = self.model(data) + loss = self.criterion(out, data.y) + + self.scaler.scale(loss).backward() + self.scaler.step(self.online_optimizer) + self.scaler.update() + + # Memory management + self.memory_manager.auto_manage_memory(threshold=0.8) + + return loss.item() + + except Exception as e: + logger.error(f"Error during online learning: {str(e)}", exc_info=True) + self.memory_manager.empty_cache() + return None + + def save_model(self, path: str = None): + """Save the model weights with AMD-specific optimizations""" + if path is None: + path = os.path.join(config.MODEL_DIR, f"{config.MODEL_NAME}.pt") + + # Save model state with additional information + checkpoint = { + "model_state_dict": self.model.state_dict(), + "optimizer_state_dict": self.optimizer.state_dict(), + "scheduler_state_dict": self.scheduler.state_dict(), + "scaler_state_dict": self.scaler.state_dict(), + "train_losses": self.train_losses, + "val_losses": self.val_losses, + "best_val_loss": self.best_val_loss, + "config": { + "device": config.DEVICE, + "mixed_precision": config.MIXED_PRECISION, + "precision": config.PRECISION, + "rocm_opt_level": config.ROCM_OPT_LEVEL if config.AMD_GPU else None, + "batch_size": config.BATCH_SIZE, + "hidden_channels": config.HIDDEN_CHANNELS, + "num_heads": config.NUM_HEADS, + }, + } + + torch.save(checkpoint, path) + logger.info(f"Model saved to {path}") + + def load_model(self, path: str = None): + """Load model weights with AMD-specific optimizations""" + if path is None: + path = os.path.join(config.MODEL_DIR, f"{config.MODEL_NAME}.pt") + + if os.path.exists(path): + try: + # Check memory before loading model + if not self.memory_manager.ensure_memory(4 * 1024**3): # 4GB + logger.warning( + "Not enough memory to load model. Falling back to CPU." + ) + self.device = torch.device("cpu") + self.model = self.model.to(self.device) + + checkpoint = torch.load(path, map_location=self.device) + + self.model.load_state_dict(checkpoint["model_state_dict"]) + self.optimizer.load_state_dict(checkpoint["optimizer_state_dict"]) + self.scheduler.load_state_dict(checkpoint["scheduler_state_dict"]) + self.scaler.load_state_dict(checkpoint["scaler_state_dict"]) + + self.train_losses = checkpoint.get("train_losses", []) + self.val_losses = checkpoint.get("val_losses", []) + self.best_val_loss = checkpoint.get("best_val_loss", float("inf")) + + # Restore config if available + if "config" in checkpoint: + saved_config = checkpoint["config"] + if saved_config.get("device") != config.DEVICE: + logger.warning( + f"Model was trained on {saved_config['device']} but current device is {config.DEVICE}" + ) + + logger.info(f"Model loaded from {path}") + + # Re-optimize model for current hardware + self.model = self.amd_optimizer.optimize_model(self.model) + + return True + + except Exception as e: + logger.error(f"Error loading model: {str(e)}", exc_info=True) + return False + else: + logger.warning(f"Model file not found at {path}") + return False + + def predict(self, data, return_attention: bool = False): + """Make predictions on new data with AMD optimizations""" + self.model.eval() + with torch.no_grad(): + # Check memory before prediction + if not self.memory_manager.ensure_memory(1 * 1024**3): # 1GB + logger.warning("Prediction skipped due to memory constraints") + return None + + try: + data = data.to(self.device, non_blocking=config.PIN_MEMORY) + + with autocast( + enabled=config.MIXED_PRECISION, + dtype=self.amd_optimizer.get_precision_dtype(), + ): + if return_attention: + out = self.model(data) + attention_weights = self.model.get_attention_weights(data) + return out, attention_weights + else: + return self.model(data) + + except Exception as e: + logger.error(f"Error during prediction: {str(e)}", exc_info=True) + self.memory_manager.empty_cache() + return None + + def benchmark(self, input_data, num_runs: int = 100): + """Benchmark model performance on AMD GPU""" + return self.amd_optimizer.benchmark_model(self.model, input_data, num_runs) diff --git a/src/trading/__init__.py b/src/trading/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/trading/broker.py b/src/trading/broker.py new file mode 100644 index 0000000..154da54 --- /dev/null +++ b/src/trading/broker.py @@ -0,0 +1,38 @@ +""" +Base broker interface for order execution. +""" + +import logging +from abc import ABC, abstractmethod +from typing import Dict, Optional + +logger = logging.getLogger(__name__) + + +class Broker(ABC): + """Abstract base class for broker implementations.""" + + @abstractmethod + def submit_order(self, order: Dict) -> Optional[str]: + """Submit an order and return the order ID if accepted.""" + pass + + @abstractmethod + def cancel_order(self, order_id: str) -> bool: + """Cancel an existing order.""" + pass + + @abstractmethod + def get_order_status(self, order_id: str) -> Dict: + """Get the status of an order.""" + pass + + @abstractmethod + def get_positions(self) -> Dict: + """Get current positions.""" + pass + + @abstractmethod + def get_account_summary(self) -> Dict: + """Get account summary.""" + pass diff --git a/src/trading/ib_broker.py b/src/trading/ib_broker.py new file mode 100644 index 0000000..30c6087 --- /dev/null +++ b/src/trading/ib_broker.py @@ -0,0 +1,47 @@ +""" +Interactive Brokers broker implementation. +""" + +import logging +from typing import Dict, Optional + +from src.trading.broker import Broker + +logger = logging.getLogger(__name__) + + +class InteractiveBrokersBroker(Broker): + """Broker implementation for Interactive Brokers.""" + + def __init__(self, host: str = "127.0.0.1", port: int = 7497, client_id: int = 1): + self.host = host + self.port = port + self.client_id = client_id + # Placeholder: initialize ib_insync connection + + def submit_order(self, order: Dict) -> Optional[str]: + """Submit an order via Interactive Brokers.""" + logger.info(f"Submitting order via IB: {order}") + # Placeholder: implement IB order submission + return None + + def cancel_order(self, order_id: str) -> bool: + """Cancel an order via Interactive Brokers.""" + logger.info(f"Cancelling order {order_id} via IB") + # Placeholder: implement IB order cancellation + return False + + def get_order_status(self, order_id: str) -> Dict: + """Get order status from Interactive Brokers.""" + # Placeholder: implement IB order status retrieval + return {} + + def get_positions(self) -> Dict: + """Get positions from Interactive Brokers.""" + # Placeholder: implement IB positions retrieval + return {} + + def get_account_summary(self) -> Dict: + """Get account summary from Interactive Brokers.""" + # Placeholder: implement IB account summary retrieval + return {} diff --git a/src/trading/paper_broker.py b/src/trading/paper_broker.py new file mode 100644 index 0000000..cea115a --- /dev/null +++ b/src/trading/paper_broker.py @@ -0,0 +1,82 @@ +""" +Paper trading broker implementation for backtesting and simulation. +""" + +import logging +import uuid +from typing import Dict, Optional + +from config import config +from src.trading.broker import Broker + +logger = logging.getLogger(__name__) + + +class PaperTradingBroker(Broker): + """Simulated broker for paper trading.""" + + def __init__(self, initial_cash: float = 100000.0): + self.cash = initial_cash + self.positions = {} + self.orders = {} + self.transaction_cost = config.TRANSACTION_COST + + def submit_order(self, order: Dict) -> Optional[str]: + """Submit a simulated order.""" + order_id = str(uuid.uuid4()) + self.orders[order_id] = {**order, "status": "filled"} + + ticker = order["ticker"] + quantity = order["quantity"] + price = order["price"] + action = order["action"] + cost = quantity * price * (1 + self.transaction_cost) + + if action == "buy": + if cost > self.cash: + logger.warning(f"Insufficient cash for buy order: {order_id}") + self.orders[order_id]["status"] = "rejected" + return None + self.cash -= cost + self.positions[ticker] = self.positions.get(ticker, 0) + quantity + elif action == "sell": + if self.positions.get(ticker, 0) < quantity: + logger.warning(f"Insufficient shares for sell order: {order_id}") + self.orders[order_id]["status"] = "rejected" + return None + self.cash += quantity * price * (1 - self.transaction_cost) + self.positions[ticker] -= quantity + if self.positions[ticker] == 0: + del self.positions[ticker] + + logger.info( + f"Paper order filled: {order_id} - {action} {quantity} {ticker} @ {price}" + ) + return order_id + + def cancel_order(self, order_id: str) -> bool: + """Cancel a simulated order.""" + if order_id in self.orders: + self.orders[order_id]["status"] = "cancelled" + return True + return False + + def get_order_status(self, order_id: str) -> Dict: + """Get the status of a simulated order.""" + return self.orders.get(order_id, {}) + + def get_positions(self) -> Dict: + """Get current simulated positions.""" + return self.positions.copy() + + def get_account_summary(self) -> Dict: + """Get simulated account summary.""" + total_value = self.cash + sum( + self.positions.get(t, 0) * 100 # placeholder price + for t in self.positions + ) + return { + "cash": self.cash, + "positions": self.positions.copy(), + "total_value": total_value, + } diff --git a/src/trading/real_time_trader.py b/src/trading/real_time_trader.py new file mode 100644 index 0000000..7b98a1f --- /dev/null +++ b/src/trading/real_time_trader.py @@ -0,0 +1,91 @@ +""" +Real-time trader for live trading execution. +""" + +import logging +from datetime import datetime +from typing import Dict + +from config import config + +logger = logging.getLogger(__name__) + + +class RealTimeTrader: + """Manage real-time trading execution and risk.""" + + def __init__(self, model, pipeline, broker): + self.model = model + self.pipeline = pipeline + self.broker = broker + self.current_positions = {} + self.pending_orders = {} + self.daily_pnl = 0.0 + self.max_drawdown = 0.0 + self.entry_times = {} + + def _calculate_position_size(self, ticker: str, price: float) -> int: + """Calculate position size based on risk management rules.""" + account_summary = self.broker.get_account_summary() + total_value = account_summary.get("total_value", config.INITIAL_CAPITAL) + max_position_value = total_value * config.MAX_POSITION_SIZE + position_size = int(max_position_value / price) + return max(0, position_size) + + def _check_holding_period(self, ticker: str, timestamp: str) -> bool: + """Check if position holding period constraints are met.""" + if ticker not in self.entry_times: + return False + entry_time = datetime.strptime(self.entry_times[ticker], "%Y-%m-%d %H:%M:%S") + current_time = datetime.strptime(timestamp, "%Y-%m-%d %H:%M:%S") + # Placeholder: implement holding period checks + return False + + def _process_pending_orders(self): + """Process and update pending orders.""" + for order_id, order in list(self.pending_orders.items()): + status = self.broker.get_order_status(order_id) + if status.get("status") == "filled": + ticker = order["ticker"] + action = order["action"] + quantity = order["quantity"] + if action == "buy": + self.current_positions[ticker] = ( + self.current_positions.get(ticker, 0) + quantity + ) + self.entry_times[ticker] = order["timestamp"] + elif action == "sell": + self.current_positions[ticker] = ( + self.current_positions.get(ticker, 0) - quantity + ) + if self.current_positions[ticker] <= 0: + del self.current_positions[ticker] + del self.entry_times[ticker] + del self.pending_orders[order_id] + logger.info(f"Order {order_id} processed: {action} {quantity} {ticker}") + + def _update_portfolio_value(self): + """Update portfolio value and track P&L.""" + # Placeholder: implement portfolio value tracking + pass + + def _check_risk_limits(self) -> bool: + """Check if risk limits have been exceeded.""" + # Placeholder: implement risk limit checks + return False + + def _close_all_positions(self): + """Close all open positions.""" + for ticker in list(self.current_positions.keys()): + order = { + "ticker": ticker, + "action": "sell", + "quantity": self.current_positions[ticker], + "price": 0, # market order + "timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"), + "type": "market", + } + order_id = self.broker.submit_order(order) + if order_id: + self.pending_orders[order_id] = order + logger.info(f"Submitted close order for {ticker}") diff --git a/src/utils/__init__.py b/src/utils/__init__.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/src/utils/__init__.py @@ -0,0 +1 @@ + diff --git a/src/utils/helpers.py b/src/utils/helpers.py new file mode 100644 index 0000000..688fb08 --- /dev/null +++ b/src/utils/helpers.py @@ -0,0 +1,46 @@ +""" +Helper functions for the trading GNN project. +""" + +from datetime import datetime, timedelta +from typing import List + +from config import config + + +def generate_intraday_timestamps(date: str) -> List[str]: + """ + Generate all intraday timestamps for a given trading date. + + Parameters: + date: Trading date string (YYYY-MM-DD). + + Returns: + List of timestamp strings for the trading day. + """ + market_open = datetime.strptime(config.TRADING_HOURS["start"], "%H:%M").time() + market_close = datetime.strptime(config.TRADING_HOURS["end"], "%H:%M").time() + + open_dt = datetime.strptime(f"{date} {market_open}", "%Y-%m-%d %H:%M:%S") + close_dt = datetime.strptime(f"{date} {market_close}", "%Y-%m-%d %H:%M:%S") + + if config.TRADING_FREQUENCY == "1min": + delta = timedelta(minutes=1) + elif config.TRADING_FREQUENCY == "5min": + delta = timedelta(minutes=5) + elif config.TRADING_FREQUENCY == "15min": + delta = timedelta(minutes=15) + elif config.TRADING_FREQUENCY == "30min": + delta = timedelta(minutes=30) + elif config.TRADING_FREQUENCY == "1h": + delta = timedelta(hours=1) + else: + delta = timedelta(minutes=1) + + timestamps = [] + current = open_dt + while current <= close_dt: + timestamps.append(current.strftime("%Y-%m-%d %H:%M:%S")) + current += delta + + return timestamps diff --git a/src/utils/memory_manager.py b/src/utils/memory_manager.py new file mode 100644 index 0000000..534e04d --- /dev/null +++ b/src/utils/memory_manager.py @@ -0,0 +1,215 @@ +import gc +import logging +from typing import Any, Dict, Optional + +import torch + +from config import config + +logger = logging.getLogger(__name__) + + +class MemoryManager: + """ + Memory management for AMD Radeon R9700 AI Pro (32GB) + """ + + def __init__(self): + self.device = torch.device(config.DEVICE) + self.max_memory = 0 + self.memory_limit = 0 + self._initialize_memory() + + def _initialize_memory(self): + """Initialize memory settings""" + if config.DEVICE == "cuda": + try: + # Get total GPU memory + self.max_memory = torch.cuda.get_device_properties(0).total_memory + self.memory_limit = int(self.max_memory * config.GPU_MEMORY_LIMIT) + + # Set memory limit + torch.cuda.set_per_process_memory_fraction(config.GPU_MEMORY_LIMIT, 0) + + logger.info(f"Initialized memory manager for AMD Radeon R9700 AI Pro") + logger.info(f"Total GPU Memory: {self.max_memory / 1024**3:.2f}GB") + logger.info( + f"Memory Limit: {self.memory_limit / 1024**3:.2f}GB ({config.GPU_MEMORY_LIMIT * 100:.0f}%)" + ) + + except Exception as e: + logger.error(f"Error initializing memory manager: {str(e)}") + self.max_memory = 0 + self.memory_limit = 0 + + def empty_cache(self): + """Clear the GPU cache and run garbage collection""" + if config.DEVICE == "cuda": + torch.cuda.empty_cache() + gc.collect() + + def check_memory(self) -> Dict[str, Any]: + """Check current GPU memory usage""" + if config.DEVICE != "cuda": + return { + "allocated": 0, + "max_allocated": 0, + "total": 0, + "limit": 0, + "usage_percent": 0, + "free": 0, + } + + try: + allocated = torch.cuda.memory_allocated(0) + max_allocated = torch.cuda.max_memory_allocated(0) + free = self.memory_limit - allocated + + return { + "allocated": allocated, + "max_allocated": max_allocated, + "total": self.max_memory, + "limit": self.memory_limit, + "usage_percent": (allocated / self.memory_limit) * 100, + "free": free, + } + except Exception as e: + logger.error(f"Error checking memory: {str(e)}") + return { + "allocated": 0, + "max_allocated": 0, + "total": 0, + "limit": 0, + "usage_percent": 0, + "free": 0, + } + + def ensure_memory(self, required_memory: int) -> bool: + """ + Ensure there's enough memory for an operation + + Parameters: + required_memory: Memory required in bytes + + Returns: + True if there's enough memory, False otherwise + """ + if config.DEVICE != "cuda": + return True + + memory_info = self.check_memory() + if memory_info["allocated"] + required_memory > memory_info["limit"]: + # Try to free memory + self.empty_cache() + memory_info = self.check_memory() + + if memory_info["allocated"] + required_memory > memory_info["limit"]: + logger.warning( + f"Not enough GPU memory. Required: {required_memory / 1024**2:.2f}MB, " + f"Available: {memory_info['free'] / 1024**2:.2f}MB" + ) + return False + + return True + + def auto_manage_memory(self, threshold: float = 0.85): + """ + Automatically manage memory based on usage + + Parameters: + threshold: Memory usage threshold (0-1) to trigger cleanup + """ + if config.DEVICE != "cuda": + return + + memory_info = self.check_memory() + if memory_info["usage_percent"] > threshold * 100: + logger.info( + f"High GPU memory usage: {memory_info['usage_percent']:.2f}%. Clearing cache." + ) + self.empty_cache() + + def get_memory_stats(self) -> str: + """Get formatted memory statistics""" + memory_info = self.check_memory() + return ( + f"GPU Memory Usage:\n" + f" Allocated: {memory_info['allocated'] / 1024**3:.2f}GB\n" + f" Max Allocated: {memory_info['max_allocated'] / 1024**3:.2f}GB\n" + f" Total: {memory_info['total'] / 1024**3:.2f}GB\n" + f" Limit: {memory_info['limit'] / 1024**3:.2f}GB\n" + f" Usage: {memory_info['usage_percent']:.2f}%\n" + f" Free: {memory_info['free'] / 1024**3:.2f}GB" + ) + + def estimate_model_memory(self, model: torch.nn.Module) -> int: + """ + Estimate memory required for a model + + Parameters: + model: PyTorch model + + Returns: + Estimated memory in bytes + """ + if config.DEVICE != "cuda": + return 0 + + try: + # Move model to GPU to get accurate memory estimate + model = model.to(config.DEVICE) + + # Get model parameters + param_size = 0 + for param in model.parameters(): + param_size += param.nelement() * param.element_size() + + # Get model buffers + buffer_size = 0 + for buffer in model.buffers(): + buffer_size += buffer.nelement() * buffer.element_size() + + # Estimate forward pass memory (activations) + # This is a rough estimate - actual memory usage may vary + forward_memory = ( + param_size * 2 + ) # Activations typically use 2x parameter memory + + # Total memory estimate + total_memory = param_size + buffer_size + forward_memory + + # Add some buffer for overhead + total_memory = int(total_memory * 1.2) + + return total_memory + + except Exception as e: + logger.error(f"Error estimating model memory: {str(e)}") + return 0 + + def log_memory_usage(self, tag: str = ""): + """Log current memory usage""" + memory_info = self.check_memory() + logger.info( + f"Memory Usage {tag}: " + f"Allocated={memory_info['allocated'] / 1024**3:.2f}GB, " + f"Usage={memory_info['usage_percent']:.2f}%" + ) + + def monitor_memory(self, interval: float = 60.0): + """ + Monitor memory usage in a background thread + + Parameters: + interval: Monitoring interval in seconds + """ + import threading + + def monitor(): + while True: + self.auto_manage_memory() + self.log_memory_usage("[Monitor]") + time.sleep(interval) + + thread = threading.Thread(target=monitor, daemon=True) + thread.start() diff --git a/src/utils/text_processing.py b/src/utils/text_processing.py new file mode 100644 index 0000000..ea28459 --- /dev/null +++ b/src/utils/text_processing.py @@ -0,0 +1,12 @@ +""" +Text processing utilities for sentiment and news analysis. +""" + +import re + + +def clean_text(text: str) -> str: + """Clean and normalize text.""" + text = re.sub(r"http\S+", "", text) + text = re.sub(r"[^\w\s]", "", text) + return text.strip().lower() diff --git a/src/utils/visualization.py b/src/utils/visualization.py new file mode 100644 index 0000000..ec440b4 --- /dev/null +++ b/src/utils/visualization.py @@ -0,0 +1,39 @@ +""" +Visualization tools for the trading GNN project. +""" + +import matplotlib.pyplot as plt +import pandas as pd + + +def plot_performance(portfolio_values, benchmark_values, filename): + """Plot portfolio performance against benchmark.""" + plt.figure(figsize=(12, 6)) + plt.plot(portfolio_values.index, portfolio_values.values, label="Portfolio") + plt.plot(benchmark_values.index, benchmark_values.values, label="Benchmark") + plt.title("Portfolio vs Benchmark Performance") + plt.xlabel("Date") + plt.ylabel("Value") + plt.legend() + plt.savefig(filename) + plt.close() + + +def plot_trade_log(trade_log, filename): + """Plot trade log entries.""" + # Placeholder implementation + plt.figure(figsize=(12, 6)) + plt.title("Trade Log") + plt.xlabel("Date") + plt.ylabel("Trade") + plt.savefig(filename) + plt.close() + + +def plot_feature_importance(feature_importance, filename): + """Plot feature importance.""" + # Placeholder implementation + plt.figure(figsize=(10, 6)) + plt.title("Feature Importance") + plt.savefig(filename) + plt.close() diff --git a/src/web/__init__.py b/src/web/__init__.py new file mode 100644 index 0000000..9b0a6e1 --- /dev/null +++ b/src/web/__init__.py @@ -0,0 +1 @@ +# Web frontend package diff --git a/src/web/api/__init__.py b/src/web/api/__init__.py new file mode 100644 index 0000000..3513a80 --- /dev/null +++ b/src/web/api/__init__.py @@ -0,0 +1 @@ +# web api package diff --git a/src/web/api/dashboard.py b/src/web/api/dashboard.py new file mode 100644 index 0000000..7ec7823 --- /dev/null +++ b/src/web/api/dashboard.py @@ -0,0 +1,84 @@ +""" +Dashboard API endpoints. +""" + +import logging +from typing import Dict + +from fastapi import APIRouter, HTTPException + +from config import config +from src.web.services.state import AppState + +logger = logging.getLogger(__name__) +router = APIRouter() + +# Reference to global app state (injected via module import in app.py) +app_state: AppState = None # type: ignore + + +def _get_state() -> AppState: + from src.web.app import app_state as _state + + return _state + + +@router.get("/metrics") +async def get_metrics() -> Dict: + """Get current dashboard metrics (memory, account, model).""" + state = _get_state() + mm = state.memory_manager + broker = state.broker + + memory = mm.check_memory() if mm else {} + account = broker.get_account_summary() if broker else {} + + return { + "memory": { + "allocated_gb": round(memory.get("allocated", 0) / 1024**3, 2), + "max_allocated_gb": round(memory.get("max_allocated", 0) / 1024**3, 2), + "total_gb": round(memory.get("total", 0) / 1024**3, 2), + "limit_gb": round(memory.get("limit", 0) / 1024**3, 2), + "usage_percent": round(memory.get("usage_percent", 0), 1), + "free_gb": round(memory.get("free", 0) / 1024**3, 2), + }, + "account": { + "cash": round(account.get("cash", 0), 2), + "total_value": round(account.get("total_value", 0), 2), + "positions_count": len(account.get("positions", {})), + "positions": account.get("positions", {}), + }, + "model": { + "status": state.model_status, + "device": config.DEVICE, + "amd_gpu": config.AMD_GPU, + "mixed_precision": config.MIXED_PRECISION, + "precision": config.PRECISION, + "hidden_channels": config.HIDDEN_CHANNELS, + "num_heads": config.NUM_HEADS, + "batch_size": config.BATCH_SIZE, + "learning_rate": config.LEARNING_RATE, + }, + "system": { + "project_name": config.PROJECT_NAME, + "version": config.VERSION, + "training_active": state.training_active, + "trading_active": state.trading_active, + }, + } + + +@router.get("/logs") +async def get_logs(limit: int = 100) -> Dict: + """Get recent log entries.""" + state = _get_state() + logs = state.logs[-limit:] if state.logs else [] + return {"logs": logs, "total": len(state.logs)} + + +@router.post("/logs/clear") +async def clear_logs() -> Dict: + """Clear stored log entries.""" + state = _get_state() + state.logs.clear() + return {"status": "cleared"} diff --git a/src/web/api/data_endpoints.py b/src/web/api/data_endpoints.py new file mode 100644 index 0000000..3238d28 --- /dev/null +++ b/src/web/api/data_endpoints.py @@ -0,0 +1,139 @@ +""" +Data pipeline API endpoints. +""" + +import logging +from datetime import datetime +from typing import Dict, List + +from fastapi import APIRouter + +from config import config + +logger = logging.getLogger(__name__) +router = APIRouter() + + +def _get_state(): + from src.web.app import app_state as _state + + return _state + + +@router.get("/tickers") +async def get_tickers() -> Dict: + """Get the list of tracked tickers.""" + return { + "initial": config.INITIAL_TICKERS, + "index": config.INDEX_TICKER, + "count": len(config.INITIAL_TICKERS), + } + + +@router.get("/pipeline/status") +async def get_pipeline_status() -> Dict: + """Get data pipeline status.""" + state = _get_state() + pipeline = state.pipeline + if not pipeline: + return {"status": "not_initialized"} + + return { + "status": "ready", + "tickers_loaded": len(pipeline.price_data), + "db_path": pipeline.db_path, + } + + +@router.post("/update") +async def update_data() -> Dict: + """Trigger a data pipeline update.""" + state = _get_state() + pipeline = state.pipeline + if not pipeline: + return {"status": "error", "message": "Pipeline not initialized"} + + import asyncio + + async def _update(): + try: + state.add_log(f"[{datetime.now()}] Data update started") + pipeline.update_all_data() + state.add_log(f"[{datetime.now()}] Data update completed") + except Exception as e: + state.add_log(f"[{datetime.now()}] Data update error: {e}") + logger.error(f"Data update error: {e}", exc_info=True) + + asyncio.create_task(_update()) + return {"status": "started"} + + +@router.get("/features/{ticker}") +async def get_features(ticker: str) -> Dict: + """Get latest features for a ticker.""" + state = _get_state() + pipeline = state.pipeline + if not pipeline: + return {"status": "error", "message": "Pipeline not initialized"} + + timestamp = datetime.now().strftime("%Y-%m-%d %H:%M:%S") + features = pipeline.get_latest_features([ticker], timestamp) + return {"ticker": ticker, "timestamp": timestamp, "features": features} + + +@router.get("/price/{ticker}") +async def get_price(ticker: str) -> Dict: + """Get latest price data for a ticker.""" + state = _get_state() + pipeline = state.pipeline + if not pipeline or ticker not in pipeline.price_data: + return {"status": "error", "message": "Ticker not found"} + + df = pipeline.price_data[ticker] + if df.empty: + return {"status": "error", "message": "No data available"} + + latest = df.iloc[-1] + return { + "ticker": ticker, + "date": str(df.index[-1]), + "open": latest["Open"], + "high": latest["High"], + "low": latest["Low"], + "close": latest["Close"], + "adj_close": latest["Adj Close"], + "volume": int(latest["Volume"]), + } + + +@router.get("/prices/{ticker}") +async def get_price_history(ticker: str, limit: int = 30) -> List[Dict]: + """Get historical price data for a ticker.""" + state = _get_state() + pipeline = state.pipeline + if not pipeline or ticker not in pipeline.price_data: + return [] + + df = pipeline.price_data[ticker].tail(limit) + return [ + { + "date": str(idx), + "open": row["Open"], + "high": row["High"], + "low": row["Low"], + "close": row["Close"], + "volume": int(row["Volume"]), + } + for idx, row in df.iterrows() + ] + + +@router.get("/corporate-actions/{ticker}") +async def get_corporate_actions(ticker: str) -> Dict: + """Get corporate actions for a ticker.""" + state = _get_state() + pipeline = state.pipeline + if not pipeline or ticker not in pipeline.corporate_actions: + return {"ticker": ticker, "actions": {}} + + return {"ticker": ticker, "actions": pipeline.corporate_actions[ticker]} diff --git a/src/web/api/models_endpoints.py b/src/web/api/models_endpoints.py new file mode 100644 index 0000000..7922538 --- /dev/null +++ b/src/web/api/models_endpoints.py @@ -0,0 +1,193 @@ +""" +Model management API endpoints. +""" + +import logging +import time +from datetime import datetime +from typing import Dict + +from fastapi import APIRouter + +from config import config + +logger = logging.getLogger(__name__) +router = APIRouter() + + +def _get_state(): + from src.web.app import app_state as _state + + return _state + + +@router.get("/status") +async def get_model_status() -> Dict: + """Get current model status.""" + state = _get_state() + return { + "status": state.model_status, + "training_active": state.training_active, + "device": config.DEVICE, + "amd_gpu": config.AMD_GPU, + "mixed_precision": config.MIXED_PRECISION, + "precision": config.PRECISION, + "rocm_opt_level": config.ROCM_OPT_LEVEL, + "model_name": config.MODEL_NAME, + "hidden_channels": config.HIDDEN_CHANNELS, + "num_heads": config.NUM_HEADS, + "dropout": config.DROPOUT, + "learning_rate": config.LEARNING_RATE, + "batch_size": config.BATCH_SIZE, + "epochs": config.EPOCHS, + "sequence_length": config.SEQUENCE_LENGTH, + } + + +@router.post("/train") +async def start_training() -> Dict: + """Start a model training run.""" + state = _get_state() + if state.training_active: + return {"status": "already_training"} + + state.training_active = True + state.add_log(f"[{datetime.now()}] Training started") + logger.info("Training started from web frontend") + + # Run training in background (simplified: kick off in async task) + import asyncio + + async def _train(): + try: + trainer = state.trainer + pipeline = state.pipeline + train_dataset = pipeline.create_training_dataset() + val_dataset = pipeline.create_validation_dataset( + start_date=config.TRAIN_END_DATE, end_date=config.VAL_END_DATE + ) + state.add_log( + f"[{datetime.now()}] Datasets created: train={len(train_dataset)}, val={len(val_dataset)}" + ) + # Note: actual training would happen here + # trainer.train(train_dataset, val_dataset) + await asyncio.sleep(2) # placeholder for actual training + state.model_status = "trained" + state.add_log(f"[{datetime.now()}] Training completed") + except Exception as e: + state.add_log(f"[{datetime.now()}] Training error: {e}") + logger.error(f"Training error: {e}", exc_info=True) + finally: + state.training_active = False + + asyncio.create_task(_train()) + return {"status": "started"} + + +@router.post("/train/stop") +async def stop_training() -> Dict: + """Stop the current training run.""" + state = _get_state() + if not state.training_active: + return {"status": "not_training"} + state.training_active = False + state.add_log(f"[{datetime.now()}] Training stopped") + logger.info("Training stopped from web frontend") + return {"status": "stopped"} + + +@router.post("/save") +async def save_model() -> Dict: + """Save the current model weights.""" + state = _get_state() + trainer = state.trainer + if not trainer: + return {"status": "error", "message": "Trainer not initialized"} + try: + trainer.save_model() + state.model_status = "saved" + state.add_log(f"[{datetime.now()}] Model saved") + return {"status": "saved"} + except Exception as e: + return {"status": "error", "message": str(e)} + + +@router.post("/load") +async def load_model() -> Dict: + """Load model weights from disk.""" + state = _get_state() + trainer = state.trainer + if not trainer: + return {"status": "error", "message": "Trainer not initialized"} + try: + success = trainer.load_model() + state.model_status = "loaded" if success else "untrained" + state.add_log(f"[{datetime.now()}] Model loaded: {success}") + return {"status": "loaded" if success else "failed"} + except Exception as e: + return {"status": "error", "message": str(e)} + + +@router.post("/benchmark") +async def run_benchmark() -> Dict: + """Run a performance benchmark on the model.""" + state = _get_state() + model = state.model + amd_optimizer = state.amd_optimizer + memory_manager = state.memory_manager + + if not model: + return {"status": "error", "message": "Model not initialized"} + + import torch + import torch.nn as nn + + num_features = len(config.INTRADAY_FEATURES) + 5 + num_stocks = 50 + num_edges = 200 + + x = torch.randn(num_stocks, config.SEQUENCE_LENGTH, num_features).to(config.DEVICE) + edge_index = torch.randint(0, num_stocks, (2, num_edges)).to(config.DEVICE) + edge_attr = torch.randn(num_edges, 1).to(config.DEVICE) + + # Warm-up + for _ in range(10): + with torch.no_grad(): + _ = model((x, edge_index, edge_attr)) + + # Benchmark inference + start = time.time() + num_runs = 100 + for _ in range(num_runs): + with torch.no_grad(): + _ = model((x, edge_index, edge_attr)) + inference_time = (time.time() - start) / num_runs + + # Benchmark training + model.train() + optimizer = torch.optim.Adam(model.parameters(), lr=config.LEARNING_RATE) + criterion = nn.MSELoss() + y = torch.randn(num_stocks, 1).to(config.DEVICE) + + start = time.time() + for _ in range(num_runs): + optimizer.zero_grad() + out = model((x, edge_index, edge_attr)) + loss = criterion(out, y) + loss.backward() + optimizer.step() + training_time = (time.time() - start) / num_runs + + memory = memory_manager.check_memory() if memory_manager else {} + + result = { + "status": "complete", + "inference_time_ms": round(inference_time * 1000, 3), + "training_time_ms": round(training_time * 1000, 3), + "inference_throughput": round(1 / inference_time, 2), + "training_throughput": round(1 / training_time, 2), + "memory_allocated_gb": round(memory.get("allocated", 0) / 1024**3, 2), + "device": config.DEVICE, + } + state.add_log(f"[{datetime.now()}] Benchmark: {result}") + return result diff --git a/src/web/api/trading_endpoints.py b/src/web/api/trading_endpoints.py new file mode 100644 index 0000000..da9e652 --- /dev/null +++ b/src/web/api/trading_endpoints.py @@ -0,0 +1,145 @@ +""" +Trading API endpoints. +""" + +import logging +from datetime import datetime +from typing import Dict, List + +from fastapi import APIRouter + +from config import config + +logger = logging.getLogger(__name__) +router = APIRouter() + + +def _get_state(): + from src.web.app import app_state as _state + + return _state + + +@router.get("/status") +async def get_trading_status() -> Dict: + """Get current trading system status.""" + state = _get_state() + broker = state.broker + account = broker.get_account_summary() if broker else {} + return { + "active": state.trading_active, + "cash": round(account.get("cash", 0), 2), + "total_value": round(account.get("total_value", 0), 2), + "positions": account.get("positions", {}), + "orders_count": len(broker.orders) if broker else 0, + } + + +@router.post("/start") +async def start_trading() -> Dict: + """Start the live trading system.""" + state = _get_state() + if state.trading_active: + return {"status": "already_running"} + state.trading_active = True + state.add_log(f"[{datetime.now()}] Trading started") + logger.info("Trading started from web frontend") + return {"status": "started"} + + +@router.post("/stop") +async def stop_trading() -> Dict: + """Stop the live trading system.""" + state = _get_state() + if not state.trading_active: + return {"status": "not_running"} + state.trading_active = False + state.add_log(f"[{datetime.now()}] Trading stopped") + logger.info("Trading stopped from web frontend") + return {"status": "stopped"} + + +@router.get("/orders") +async def get_orders() -> List[Dict]: + """Get all orders.""" + state = _get_state() + broker = state.broker + if not broker: + return [] + return [{"order_id": oid, **details} for oid, details in broker.orders.items()] + + +@router.post("/order") +async def submit_manual_order(order: Dict) -> Dict: + """Submit a manual order via the web frontend.""" + state = _get_state() + broker = state.broker + if not broker: + return {"status": "error", "message": "Broker not initialized"} + + order_id = broker.submit_order(order) + if order_id: + state.add_log(f"[{datetime.now()}] Manual order submitted: {order}") + return {"status": "submitted", "order_id": order_id} + return {"status": "rejected"} + + +@router.post("/cancel/{order_id}") +async def cancel_order(order_id: str) -> Dict: + """Cancel an order by ID.""" + state = _get_state() + broker = state.broker + if not broker: + return {"status": "error", "message": "Broker not initialized"} + + success = broker.cancel_order(order_id) + return {"status": "cancelled" if success else "failed"} + + +@router.post("/positions/close/{ticker}") +async def close_position(ticker: str) -> Dict: + """Close a position for a given ticker.""" + state = _get_state() + broker = state.broker + if not broker or ticker not in broker.get_positions(): + return {"status": "error", "message": "No position found"} + + quantity = broker.get_positions()[ticker] + order = { + "ticker": ticker, + "action": "sell", + "quantity": quantity, + "price": 0, + "timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"), + "type": "market", + } + order_id = broker.submit_order(order) + if order_id: + return {"status": "submitted", "order_id": order_id} + return {"status": "rejected"} + + +@router.post("/positions/close-all") +async def close_all_positions() -> Dict: + """Close all open positions.""" + state = _get_state() + broker = state.broker + if not broker: + return {"status": "error", "message": "Broker not initialized"} + + positions = broker.get_positions() + results = [] + for ticker in list(positions.keys()): + order = { + "ticker": ticker, + "action": "sell", + "quantity": positions[ticker], + "price": 0, + "timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"), + "type": "market", + } + order_id = broker.submit_order(order) + results.append({"ticker": ticker, "order_id": order_id}) + + state.add_log(f"[{datetime.now()}] Closed all positions") + return {"status": "submitted", "results": results} diff --git a/src/web/app.py b/src/web/app.py new file mode 100644 index 0000000..82eb6f5 --- /dev/null +++ b/src/web/app.py @@ -0,0 +1,193 @@ +""" +FastAPI web application for the StockGNN trading system. +""" + +import asyncio +import json +import logging +import os +from contextlib import asynccontextmanager +from typing import Any, Dict, List + +from fastapi import FastAPI, Request, WebSocket, WebSocketDisconnect +from fastapi.responses import HTMLResponse +from fastapi.staticfiles import StaticFiles +from fastapi.templating import Jinja2Templates + +from config import config +from src.amd.optimizations import AMDOptimizer +from src.data.pipeline import StockDataPipeline +from src.models.intraday_gnn import IntradayGNN +from src.models.trainer import GNNTrainer +from src.trading.paper_broker import PaperTradingBroker +from src.utils.memory_manager import MemoryManager +from src.web.api import dashboard, data_endpoints, models_endpoints, trading_endpoints +from src.web.services.state import AppState + +logger = logging.getLogger(__name__) + +# Initialize global app state +app_state = AppState() + +# Build template/static paths relative to this file +BASE_DIR = os.path.dirname(os.path.abspath(__file__)) +TEMPLATES_DIR = os.path.join(BASE_DIR, "templates") +STATIC_DIR = os.path.join(BASE_DIR, "static") + +templates = Jinja2Templates(directory=TEMPLATES_DIR) + + +@asynccontextmanager +async def lifespan(app: FastAPI): + """Application lifespan handler for startup/shutdown events.""" + logger.info("Starting StockGNN Web Frontend") + + # Initialize shared state + app_state.memory_manager = MemoryManager() + app_state.amd_optimizer = AMDOptimizer() + app_state.pipeline = StockDataPipeline() + app_state.broker = PaperTradingBroker(initial_cash=config.INITIAL_CAPITAL) + + # Initialize model (but don't train yet) + num_features = len(config.INTRADAY_FEATURES) + 5 + app_state.model = IntradayGNN(num_features, config.SEQUENCE_LENGTH) + app_state.trainer = GNNTrainer(app_state.model) + + # Try to load pre-trained weights + try: + app_state.trainer.load_model() + app_state.model_status = "loaded" + except Exception: + app_state.model_status = "untrained" + + app_state.model = app_state.amd_optimizer.optimize_model(app_state.model) + + # Start background broadcast task + app_state.broadcast_task = asyncio.create_task(_broadcast_loop()) + + yield + + # Shutdown + logger.info("Shutting down StockGNN Web Frontend") + if app_state.broadcast_task: + app_state.broadcast_task.cancel() + try: + await app_state.broadcast_task + except asyncio.CancelledError: + pass + + +app = FastAPI( + title="StockGNN R9700", + description="AMD-optimized Graph Neural Network trading system dashboard", + version=config.VERSION, + lifespan=lifespan, +) + +# Serve static files +app.mount("/static", StaticFiles(directory=STATIC_DIR), name="static") + +# Include API routers +app.include_router(dashboard.router, prefix="/api/dashboard", tags=["dashboard"]) +app.include_router(trading_endpoints.router, prefix="/api/trading", tags=["trading"]) +app.include_router(models_endpoints.router, prefix="/api/models", tags=["models"]) +app.include_router(data_endpoints.router, prefix="/api/data", tags=["data"]) + + +# --------------------------------------------------------------------------- +# HTML entry point +# --------------------------------------------------------------------------- +@app.get("/", response_class=HTMLResponse) +async def get_index(request: Request): + return templates.TemplateResponse("index.html", {"request": request}) + + +# --------------------------------------------------------------------------- +# WebSocket for real-time updates +# --------------------------------------------------------------------------- +@app.websocket("/ws") +async def websocket_endpoint(ws: WebSocket): + await ws.accept() + app_state.connections.append(ws) + logger.info(f"WebSocket client connected ({len(app_state.connections)} active)") + try: + while True: + # Wait for any incoming messages (optional client commands) + msg = await ws.receive_text() + try: + data = json.loads(msg) + await _handle_ws_message(ws, data) + except json.JSONDecodeError: + await ws.send_json({"error": "Invalid JSON"}) + except WebSocketDisconnect: + logger.info("WebSocket client disconnected") + finally: + if ws in app_state.connections: + app_state.connections.remove(ws) + + +async def _handle_ws_message(ws: WebSocket, data: Dict[str, Any]): + """Handle incoming WebSocket messages.""" + action = data.get("action") + if action == "ping": + await ws.send_json({"type": "pong", "timestamp": str(datetime.now())}) + elif action == "subscribe": + channel = data.get("channel", "all") + await ws.send_json({"type": "subscribed", "channel": channel}) + else: + await ws.send_json({"type": "error", "message": f"Unknown action: {action}"}) + + +async def _broadcast_loop(): + """Background task that periodically broadcasts metrics to all WS clients.""" + from datetime import datetime + + while True: + try: + await asyncio.sleep(2) # broadcast every 2 seconds + payload = _build_broadcast_payload() + disconnected = [] + for ws in app_state.connections: + try: + await ws.send_json(payload) + except Exception: + disconnected.append(ws) + for ws in disconnected: + if ws in app_state.connections: + app_state.connections.remove(ws) + except asyncio.CancelledError: + break + except Exception as e: + logger.error(f"Broadcast loop error: {e}") + + +def _build_broadcast_payload() -> Dict[str, Any]: + """Build the real-time metrics payload.""" + from datetime import datetime + + mm = app_state.memory_manager + broker = app_state.broker + + memory = mm.check_memory() if mm else {} + account = broker.get_account_summary() if broker else {} + + return { + "type": "metrics", + "timestamp": datetime.now().isoformat(), + "memory": { + "allocated_gb": round(memory.get("allocated", 0) / 1024**3, 2), + "total_gb": round(memory.get("total", 0) / 1024**3, 2), + "usage_percent": round(memory.get("usage_percent", 0), 1), + }, + "account": { + "cash": round(account.get("cash", 0), 2), + "total_value": round(account.get("total_value", 0), 2), + "positions": len(account.get("positions", {})), + }, + "model": { + "status": app_state.model_status, + "device": config.DEVICE, + "mixed_precision": config.MIXED_PRECISION, + "precision": config.PRECISION, + }, + } diff --git a/src/web/services/__init__.py b/src/web/services/__init__.py new file mode 100644 index 0000000..8d300ec --- /dev/null +++ b/src/web/services/__init__.py @@ -0,0 +1 @@ +# web services package diff --git a/src/web/services/state.py b/src/web/services/state.py new file mode 100644 index 0000000..568b26c --- /dev/null +++ b/src/web/services/state.py @@ -0,0 +1,31 @@ +""" +Shared application state for the web frontend. +""" + +from dataclasses import dataclass, field +from typing import Any, List, Optional + + +@dataclass +class AppState: + """Singleton-ish shared state for the web application.""" + + memory_manager: Optional[Any] = None + amd_optimizer: Optional[Any] = None + pipeline: Optional[Any] = None + broker: Optional[Any] = None + model: Optional[Any] = None + trainer: Optional[Any] = None + broadcast_task: Optional[Any] = None + + connections: List[Any] = field(default_factory=list) + model_status: str = "untrained" + training_active: bool = False + trading_active: bool = False + logs: List[str] = field(default_factory=list) + + def add_log(self, message: str): + """Add a log entry (capped at 500 lines).""" + self.logs.append(message) + if len(self.logs) > 500: + self.logs = self.logs[-500:] diff --git a/src/web/static/css/style.css b/src/web/static/css/style.css new file mode 100644 index 0000000..84a1916 --- /dev/null +++ b/src/web/static/css/style.css @@ -0,0 +1,409 @@ +/* StockGNN R9700 Dashboard Styles */ + +:root { + --bg: #0b0f17; + --surface: #111827; + --surface-2: #1a2233; + --surface-3: #243044; + --text: #e6edf3; + --text-secondary: #8b96a7; + --accent: #ff6b35; + --accent-2: #ff8c5a; + --green: #22c55e; + --green-dim: #16a34a; + --red: #ef4444; + --red-dim: #dc2626; + --blue: #3b82f6; + --blue-dim: #2563eb; + --amber: #f59e0b; + --border: rgba(255,255,255,0.06); + --radius: 10px; + --shadow: 0 4px 20px rgba(0,0,0,0.35); +} + +* { margin: 0; padding: 0; box-sizing: border-box; } + +html, body { + font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, "Helvetica Neue", Arial, sans-serif; + background: var(--bg); + color: var(--text); + height: 100%; + overflow: hidden; +} + +.app { + display: flex; + height: 100vh; + overflow: hidden; +} + +/* Sidebar */ +.sidebar { + width: 220px; + background: var(--surface); + border-right: 1px solid var(--border); + display: flex; + flex-direction: column; + flex-shrink: 0; +} + +.sidebar-header { + padding: 20px; + border-bottom: 1px solid var(--border); +} + +.logo { + display: flex; + align-items: center; + gap: 10px; + font-weight: 700; + font-size: 1.1rem; + color: var(--text); +} + +.logo-icon { + font-size: 1.5rem; + color: var(--accent); +} + +.logo-badge { + font-size: 0.7rem; + background: var(--accent); + color: #fff; + padding: 2px 6px; + border-radius: 4px; + vertical-align: middle; + margin-left: 4px; +} + +.device-badge { + margin-top: 8px; + font-size: 0.75rem; + color: var(--text-secondary); + background: var(--surface-3); + padding: 3px 8px; + border-radius: 6px; + display: inline-block; +} + +.nav { + flex: 1; + padding: 12px 0; + overflow-y: auto; +} + +.nav-item { + display: flex; + align-items: center; + gap: 12px; + padding: 12px 20px; + color: var(--text-secondary); + text-decoration: none; + font-size: 0.9rem; + border-left: 3px solid transparent; + transition: all 0.15s ease; +} + +.nav-item:hover { + color: var(--text); + background: var(--surface-2); +} + +.nav-item.active { + color: var(--accent); + background: rgba(255,107,53,0.06); + border-left-color: var(--accent); +} + +.nav-icon { + font-size: 1.1rem; + width: 22px; + text-align: center; +} + +.sidebar-footer { + padding: 16px 20px; + border-top: 1px solid var(--border); + font-size: 0.75rem; + color: var(--text-secondary); +} + +.connection-status { + display: flex; + align-items: center; + gap: 8px; + margin-bottom: 8px; +} + +.status-dot { + width: 8px; + height: 8px; + border-radius: 50%; + background: var(--green); +} +.status-dot.disconnected { background: var(--red); } +.status-dot.connecting { background: var(--amber); } + +.version { + opacity: 0.5; +} + +/* Main */ +.main { + flex: 1; + display: flex; + flex-direction: column; + overflow: hidden; +} + +.topbar { + height: 60px; + background: var(--surface); + border-bottom: 1px solid var(--border); + display: flex; + align-items: center; + justify-content: space-between; + padding: 0 24px; + flex-shrink: 0; +} + +.page-title { + font-size: 1.25rem; + font-weight: 600; + color: var(--text); +} + +.topbar-actions { + display: flex; + gap: 10px; +} + +/* Page system */ +.page { + display: none; + flex: 1; + padding: 24px; + overflow-y: auto; + gap: 24px; + flex-direction: column; +} +.page.active { display: flex; } + +/* Grid */ +.grid { + display: grid; + gap: 20px; +} +.grid-2 { grid-template-columns: repeat(2, 1fr); } +.grid-3 { grid-template-columns: repeat(3, 1fr); } +.grid-4 { grid-template-columns: repeat(4, 1fr); } + +@media (max-width: 1200px) { + .grid-4 { grid-template-columns: repeat(2, 1fr); } +} +@media (max-width: 900px) { + .grid-3 { grid-template-columns: repeat(2, 1fr); } + .grid-2 { grid-template-columns: 1fr; } + .grid-4 { grid-template-columns: 1fr; } +} +@media (max-width: 700px) { + .grid-2, .grid-3, .grid-4 { grid-template-columns: 1fr; } + .sidebar { display: none; } +} + +/* Cards */ +.card { + background: var(--surface); + border: 1px solid var(--border); + border-radius: var(--radius); + padding: 20px; + box-shadow: var(--shadow); +} + +.card-header { + display: flex; + align-items: center; + justify-content: space-between; + margin-bottom: 16px; +} +.card-header h3 { + font-size: 0.95rem; + font-weight: 600; + color: var(--text); +} + +.metric-card { + text-align: center; + padding: 24px 16px; +} + +.metric-label { + font-size: 0.8rem; + color: var(--text-secondary); + text-transform: uppercase; + letter-spacing: 0.5px; + margin-bottom: 10px; +} + +.metric-value { + font-size: 1.8rem; + font-weight: 700; + color: var(--text); + margin-bottom: 6px; +} + +.metric-sub { + font-size: 0.8rem; + color: var(--text-secondary); +} + +.chart-card { + padding: 16px; +} + +.chart-container { + position: relative; + height: 220px; +} + +/* Info grid */ +.info-grid { + display: grid; + grid-template-columns: repeat(auto-fill, minmax(280px, 1fr)); + gap: 10px 20px; +} + +.info-row { + display: flex; + justify-content: space-between; + padding: 8px 0; + border-bottom: 1px solid var(--border); + font-size: 0.9rem; +} +.info-key { color: var(--text-secondary); } +.info-value { color: var(--text); font-weight: 500; } + +/* Tables */ +.table-container { overflow-x: auto; } + +.data-table { + width: 100%; + border-collapse: collapse; + font-size: 0.85rem; +} + +.data-table th, +.data-table td { + padding: 10px 12px; + text-align: left; + border-bottom: 1px solid var(--border); + color: var(--text); +} + +.data-table th { + color: var(--text-secondary); + font-weight: 500; + text-transform: uppercase; + font-size: 0.7rem; + letter-spacing: 0.5px; +} + +.data-table tbody tr:hover { + background: var(--surface-2); +} + +/* Buttons */ +.btn { + border: none; + border-radius: 6px; + padding: 8px 16px; + font-size: 0.85rem; + font-weight: 500; + cursor: pointer; + transition: all 0.15s ease; + color: #fff; +} + +.btn-primary { background: var(--blue); } +.btn-primary:hover { background: var(--blue-dim); } + +.btn-secondary { background: var(--surface-3); color: var(--text); } +.btn-secondary:hover { background: var(--surface-2); } + +.btn-danger { background: var(--red); } +.btn-danger:hover { background: var(--red-dim); } + +.btn-accent { background: var(--accent); } +.btn-accent:hover { background: var(--accent-2); } + +.btn-sm { padding: 5px 12px; font-size: 0.75rem; } + +/* Live indicator */ +.live-indicator { + display: flex; + align-items: center; + gap: 6px; + font-size: 0.75rem; + color: var(--green); +} + +.pulse { + width: 8px; + height: 8px; + background: var(--green); + border-radius: 50%; + animation: pulse 1.5s infinite; +} + +@keyframes pulse { + 0% { opacity: 1; transform: scale(1); } + 70% { opacity: 0.4; transform: scale(1.3); } + 100% { opacity: 1; transform: scale(1); } +} + +/* Logs */ +.log-container { + background: #06080e; + border-radius: var(--radius); + padding: 16px; + max-height: 500px; + overflow-y: auto; +} + +.log-output { + font-family: "SFMono-Regular", Consolas, "Liberation Mono", Menlo, monospace; + font-size: 0.82rem; + line-height: 1.6; + color: #a0aab8; + white-space: pre-wrap; + word-break: break-word; +} + +/* Ticker grid */ +.ticker-grid { + display: grid; + grid-template-columns: repeat(auto-fill, minmax(100px, 1fr)); + gap: 8px; +} + +.ticker-chip { + background: var(--surface-2); + border: 1px solid var(--border); + border-radius: 6px; + padding: 8px 12px; + font-size: 0.85rem; + font-weight: 600; + text-align: center; + color: var(--text); +} +.ticker-chip:hover { + border-color: var(--accent); + color: var(--accent); + cursor: pointer; +} + +/* Scrollbar */ +::-webkit-scrollbar { width: 6px; height: 6px; } +::-webkit-scrollbar-track { background: transparent; } +::-webkit-scrollbar-thumb { background: var(--surface-3); border-radius: 3px; } +::-webkit-scrollbar-thumb:hover { background: #3a4a60; } diff --git a/src/web/static/js/app.js b/src/web/static/js/app.js new file mode 100644 index 0000000..807790f --- /dev/null +++ b/src/web/static/js/app.js @@ -0,0 +1,419 @@ +/** + * StockGNN R9700 Dashboard Frontend + */ + +// --------------------------------------------------------------------------- +// State +// --------------------------------------------------------------------------- +const state = { + ws: null, + wsConnected: false, + currentPage: "dashboard", + charts: {}, + metricsHistory: { + memory: [], + portfolio: [], + timestamps: [], + }, + maxHistory: 60, +}; + +// --------------------------------------------------------------------------- +// Navigation +// --------------------------------------------------------------------------- +function initNavigation() { + document.querySelectorAll(".nav-item").forEach((item) => { + item.addEventListener("click", (e) => { + e.preventDefault(); + const page = item.dataset.page; + switchPage(page); + }); + }); +} + +function switchPage(page) { + state.currentPage = page; + document.querySelectorAll(".nav-item").forEach((i) => i.classList.remove("active")); + document.querySelector(`.nav-item[data-page="${page}"]`).classList.add("active"); + document.querySelectorAll(".page").forEach((p) => p.classList.remove("active")); + document.getElementById(`page-${page}`).classList.add("active"); + document.getElementById("page-title").textContent = page.charAt(0).toUpperCase() + page.slice(1); + + if (page === "logs") loadLogs(); + if (page === "data") loadDataPage(); + if (page === "trading") loadTradingStatus(); + if (page === "models") loadModelStatus(); +} + +// --------------------------------------------------------------------------- +// WebSocket +// --------------------------------------------------------------------------- +function initWebSocket() { + const protocol = window.location.protocol === "https:" ? "wss:" : "ws:"; + const wsUrl = `${protocol}//${window.location.host}/ws`; + + state.ws = new WebSocket(wsUrl); + + state.ws.onopen = () => { + state.wsConnected = true; + updateConnectionStatus(true); + }; + + state.ws.onclose = () => { + state.wsConnected = false; + updateConnectionStatus(false); + setTimeout(initWebSocket, 3000); + }; + + state.ws.onmessage = (event) => { + const msg = JSON.parse(event.data); + handleWsMessage(msg); + }; + + state.ws.onerror = (err) => { + console.error("WebSocket error:", err); + state.wsConnected = false; + updateConnectionStatus(false); + }; +} + +function updateConnectionStatus(connected) { + const el = document.getElementById("ws-status"); + const dot = el.querySelector(".status-dot"); + const text = el.querySelector("span:last-child"); + if (connected) { + dot.classList.remove("disconnected"); + text.textContent = "Connected"; + } else { + dot.classList.add("disconnected"); + text.textContent = "Disconnected"; + } +} + +function handleWsMessage(msg) { + if (msg.type === "metrics") { + updateMetrics(msg); + } else if (msg.type === "pong") { + console.log("pong", msg.timestamp); + } +} + +// --------------------------------------------------------------------------- +// Metrics & Charts +// --------------------------------------------------------------------------- +function updateMetrics(data) { + const mem = data.memory; + const acc = data.account; + const model = data.model; + + // Update metric cards + document.getElementById("metric-memory").textContent = `${mem.usage_percent}%`; + document.getElementById("metric-memory-sub").textContent = `${mem.allocated_gb} / ${mem.total_gb} GB`; + + document.getElementById("metric-portfolio").textContent = `$${acc.total_value.toLocaleString()}`; + document.getElementById("metric-portfolio-sub").textContent = `${acc.positions} positions`; + + document.getElementById("metric-cash").textContent = `$${acc.cash.toLocaleString()}`; + + document.getElementById("metric-model").textContent = model.status; + document.getElementById("metric-model-sub").textContent = `${model.device} · ${model.precision}`; + + // Update system info + document.getElementById("sys-device").textContent = model.device; + document.getElementById("sys-precision").textContent = model.precision; + + // Update device badge + const badge = document.getElementById("device-badge"); + badge.textContent = model.device.toUpperCase(); + if (model.device === "cuda") { + badge.style.background = "var(--green-dim)"; + badge.style.color = "#fff"; + } + + // Update history + state.metricsHistory.memory.push(mem.usage_percent); + state.metricsHistory.portfolio.push(acc.total_value); + state.metricsHistory.timestamps.push(new Date(data.timestamp).toLocaleTimeString()); + if (state.metricsHistory.memory.length > state.maxHistory) { + state.metricsHistory.memory.shift(); + state.metricsHistory.portfolio.shift(); + state.metricsHistory.timestamps.shift(); + } + + updateCharts(); +} + +function initCharts() { + const ctxMem = document.getElementById("memory-chart").getContext("2d"); + const ctxPort = document.getElementById("portfolio-chart").getContext("2d"); + + state.charts.memory = new Chart(ctxMem, { + type: "line", + data: { + labels: [], + datasets: [ + { + label: "Memory %", + data: [], + borderColor: "#ff6b35", + backgroundColor: "rgba(255,107,53,0.1)", + fill: true, + tension: 0.4, + pointRadius: 0, + }, + ], + }, + options: { + responsive: true, + maintainAspectRatio: false, + plugins: { legend: { display: false } }, + scales: { + x: { display: false }, + y: { + beginAtZero: true, + max: 100, + grid: { color: "rgba(255,255,255,0.04)" }, + ticks: { color: "#8b96a7", font: { size: 10 } }, + }, + }, + }, + }); + + state.charts.portfolio = new Chart(ctxPort, { + type: "line", + data: { + labels: [], + datasets: [ + { + label: "Portfolio Value", + data: [], + borderColor: "#22c55e", + backgroundColor: "rgba(34,197,94,0.1)", + fill: true, + tension: 0.4, + pointRadius: 0, + }, + ], + }, + options: { + responsive: true, + maintainAspectRatio: false, + plugins: { legend: { display: false } }, + scales: { + x: { display: false }, + y: { + grid: { color: "rgba(255,255,255,0.04)" }, + ticks: { color: "#8b96a7", font: { size: 10 } }, + }, + }, + }, + }); +} + +function updateCharts() { + if (!state.charts.memory) return; + state.charts.memory.data.labels = state.metricsHistory.timestamps; + state.charts.memory.data.datasets[0].data = state.metricsHistory.memory; + state.charts.memory.update("none"); + + state.charts.portfolio.data.labels = state.metricsHistory.timestamps; + state.charts.portfolio.data.datasets[0].data = state.metricsHistory.portfolio; + state.charts.portfolio.update("none"); +} + +// --------------------------------------------------------------------------- +// API Helpers +// --------------------------------------------------------------------------- +async function apiGet(path) { + const res = await fetch(path); + return res.json(); +} + +async function apiPost(path, body = {}) { + const res = await fetch(path, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(body), + }); + return res.json(); +} + +// --------------------------------------------------------------------------- +// Dashboard +// --------------------------------------------------------------------------- +async function loadDashboardMetrics() { + const data = await apiGet("/api/dashboard/metrics"); + if (data.model) { + document.getElementById("sys-device").textContent = data.model.device; + document.getElementById("sys-precision").textContent = data.model.precision; + document.getElementById("sys-hidden").textContent = data.model.hidden_channels; + document.getElementById("sys-heads").textContent = data.model.num_heads; + document.getElementById("sys-batch").textContent = data.model.batch_size; + document.getElementById("sys-lr").textContent = data.model.learning_rate; + } +} + +// --------------------------------------------------------------------------- +// Trading Page +// --------------------------------------------------------------------------- +async function loadTradingStatus() { + const data = await apiGet("/api/trading/status"); + document.getElementById("trading-status").textContent = data.active ? "Active" : "Stopped"; + document.getElementById("trading-status").style.color = data.active ? "var(--green)" : "var(--text)"; + document.getElementById("orders-count").textContent = data.orders_count || 0; + document.getElementById("positions-count").textContent = data.positions ? Object.keys(data.positions).length : 0; + + // Positions table + const tbody = document.querySelector("#positions-table tbody"); + tbody.innerHTML = ""; + if (data.positions) { + Object.entries(data.positions).forEach(([ticker, qty]) => { + const tr = document.createElement("tr"); + tr.innerHTML = `${ticker}${qty}`; + tbody.appendChild(tr); + }); + } + + // Orders table + const orders = await apiGet("/api/trading/orders"); + const otbody = document.querySelector("#orders-table tbody"); + otbody.innerHTML = ""; + orders.forEach((order) => { + const tr = document.createElement("tr"); + tr.innerHTML = `${order.order_id.slice(0, 8)}${order.ticker}${order.action}${order.quantity}${order.status}`; + otbody.appendChild(tr); + }); +} + +async function closePosition(ticker) { + await apiPost(`/api/trading/positions/close/${ticker}`); + loadTradingStatus(); +} + +// --------------------------------------------------------------------------- +// Models Page +// --------------------------------------------------------------------------- +async function loadModelStatus() { + const data = await apiGet("/api/models/status"); + document.getElementById("model-status-page").textContent = data.status; + document.getElementById("model-status-detail").textContent = `${data.device} · ${data.precision}`; + document.getElementById("training-status-page").textContent = data.training_active ? "Active" : "Idle"; + document.getElementById("training-status-page").style.color = data.training_active ? "var(--green)" : "var(--text)"; +} + +// --------------------------------------------------------------------------- +// Data Page +// --------------------------------------------------------------------------- +async function loadDataPage() { + const tickers = await apiGet("/api/data/tickers"); + document.getElementById("data-tickers-count").textContent = tickers.count; + + const grid = document.getElementById("ticker-grid"); + grid.innerHTML = ""; + tickers.initial.forEach((ticker) => { + const chip = document.createElement("div"); + chip.className = "ticker-chip"; + chip.textContent = ticker; + grid.appendChild(chip); + }); + + const pipe = await apiGet("/api/data/pipeline/status"); + document.getElementById("data-pipeline-status").textContent = pipe.status; +} + +// --------------------------------------------------------------------------- +// Logs Page +// --------------------------------------------------------------------------- +async function loadLogs() { + const data = await apiGet("/api/dashboard/logs?limit=200"); + const el = document.getElementById("log-output"); + if (data.logs && data.logs.length) { + el.textContent = data.logs.join("\n"); + } else { + el.textContent = "No logs yet."; + } + // Auto-scroll + const container = document.querySelector(".log-container"); + container.scrollTop = container.scrollHeight; +} + +// --------------------------------------------------------------------------- +// Event Listeners +// --------------------------------------------------------------------------- +function initEventListeners() { + document.getElementById("refresh-btn").addEventListener("click", () => { + if (state.currentPage === "dashboard") loadDashboardMetrics(); + if (state.currentPage === "trading") loadTradingStatus(); + if (state.currentPage === "models") loadModelStatus(); + if (state.currentPage === "data") loadDataPage(); + if (state.currentPage === "logs") loadLogs(); + }); + + // Trading controls + document.getElementById("btn-start-trading").addEventListener("click", async () => { + await apiPost("/api/trading/start"); + loadTradingStatus(); + }); + document.getElementById("btn-stop-trading").addEventListener("click", async () => { + await apiPost("/api/trading/stop"); + loadTradingStatus(); + }); + document.getElementById("btn-close-all").addEventListener("click", async () => { + await apiPost("/api/trading/positions/close-all"); + loadTradingStatus(); + }); + + // Model controls + document.getElementById("btn-train").addEventListener("click", async () => { + await apiPost("/api/models/train"); + loadModelStatus(); + }); + document.getElementById("btn-stop-train").addEventListener("click", async () => { + await apiPost("/api/models/train/stop"); + loadModelStatus(); + }); + document.getElementById("btn-save-model").addEventListener("click", async () => { + const res = await apiPost("/api/models/save"); + alert(res.status); + }); + document.getElementById("btn-load-model").addEventListener("click", async () => { + const res = await apiPost("/api/models/load"); + alert(res.status); + loadModelStatus(); + }); + document.getElementById("btn-benchmark").addEventListener("click", async () => { + document.getElementById("bench-inf").textContent = "Running..."; + const res = await apiPost("/api/models/benchmark"); + if (res.status === "complete") { + document.getElementById("bench-inf").textContent = `${res.inference_time_ms} ms`; + document.getElementById("bench-train").textContent = `${res.training_time_ms} ms`; + document.getElementById("bench-inf-tput").textContent = `${res.inference_throughput} samples/s`; + document.getElementById("bench-train-tput").textContent = `${res.training_throughput} samples/s`; + document.getElementById("bench-mem").textContent = `${res.memory_allocated_gb} GB`; + document.getElementById("bench-device").textContent = res.device; + } + }); + + // Data controls + document.getElementById("btn-update-data").addEventListener("click", async () => { + await apiPost("/api/data/update"); + alert("Data update started"); + }); + + // Logs controls + document.getElementById("btn-clear-logs").addEventListener("click", async () => { + await apiPost("/api/dashboard/logs/clear"); + loadLogs(); + }); +} + +// --------------------------------------------------------------------------- +// Init +// --------------------------------------------------------------------------- +document.addEventListener("DOMContentLoaded", () => { + initNavigation(); + initWebSocket(); + initCharts(); + initEventListeners(); + loadDashboardMetrics(); +}); diff --git a/src/web/templates/index.html b/src/web/templates/index.html new file mode 100644 index 0000000..ad6e904 --- /dev/null +++ b/src/web/templates/index.html @@ -0,0 +1,268 @@ + + + + + + StockGNN R9700 Dashboard + + + + +
+ + + + +
+ +
+

Dashboard

+
+ +
+
+ + +
+
+
+
GPU Memory
+
--
+
-- / -- GB
+
+
+
Portfolio Value
+
--
+
-- positions
+
+
+
Cash
+
--
+
available
+
+
+
Model Status
+
--
+
--
+
+
+ +
+
+
+

Memory Usage

+ Live +
+
+ +
+
+
+
+

Portfolio Value

+ Live +
+
+ +
+
+
+ +
+
+

System Info

+
+
+
ProjectStockGNN R9700
+
Device--
+
Precision--
+
Hidden Channels--
+
Attention Heads--
+
Batch Size--
+
Learning Rate--
+
+
+
+ + +
+
+
+
Trading
+
Stopped
+
+ + +
+
+
+
Open Orders
+
0
+
total orders
+
+
+
Positions
+
0
+
active positions
+
+
+ +
+
+
+

Positions

+ +
+
+ + + + + +
TickerQtyAction
+
+
+
+
+

Orders

+
+
+ + + + + +
IDTickerActionQtyStatus
+
+
+
+
+ + +
+
+
+
Model Status
+
--
+
--
+
+
+
Training
+
--
+
+ + +
+
+
+
Actions
+
+ + +
+
+ +
+
+
+ +
+
+

Benchmark Results

+
+
+
Inference Time--
+
Training Time--
+
Inf Throughput--
+
Train Throughput--
+
Memory Used--
+
Device--
+
+
+
+ + +
+
+
+
Tickers
+
--
+
tracked stocks
+
+
+
Data Pipeline
+
--
+
+ +
+
+
+
Latest
+
--
+
price data
+
+
+ +
+
+

Ticker List

+
+
+
+
+ + +
+
+
+

System Logs

+ +
+
+
Loading logs...
+
+
+
+
+
+ + + +