Files
trading_gnn/live_trading.py
T
2026-05-26 13:51:02 +02:00

570 lines
21 KiB
Python

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())