570 lines
21 KiB
Python
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())
|