Tune model config for price-only training and add walk-forward CV
Downsize GNN hidden channels/heads (128→64, 16→4) to match the 30-stock price-only universe before alternative data processors are ready. Add USE_ALTERNATIVE_DATA flag (default False) that skips news/social branches so the model trains on real data rather than zero-filled stubs. Standardize target returns per-batch in the data pipeline to stabilize training. Introduce walk-forward cross-validation with expanding windows: configurable fold count and out-of-sample years. Add IC loss weighting (0.7) to complement MSE, and wire it into the trainer alongside the new loss function.
This commit is contained in:
@@ -85,14 +85,16 @@ class Config:
|
|||||||
|
|
||||||
# Model settings
|
# Model settings
|
||||||
MODEL_NAME = "stock_gnn_r9700"
|
MODEL_NAME = "stock_gnn_r9700"
|
||||||
HIDDEN_CHANNELS = 128 # Increased for R9700's compute power
|
# 64 hidden / 4 heads gives head_dim=16; appropriate for price-only 30-stock universe.
|
||||||
NUM_HEADS = 16 # Increased number of attention heads
|
# Increase once alternative data processors are implemented (see USE_ALTERNATIVE_DATA).
|
||||||
DROPOUT = 0.3 # Reduced dropout for better GPU utilization
|
HIDDEN_CHANNELS = 64
|
||||||
LEARNING_RATE = 0.0005 # Lower learning rate for stability
|
NUM_HEADS = 4
|
||||||
EPOCHS = 200 # More epochs with larger batch sizes
|
DROPOUT = 0.3
|
||||||
BATCH_SIZE = 128 # Larger batch size for R9700's memory
|
LEARNING_RATE = 0.0005
|
||||||
SEQUENCE_LENGTH = 60 # Longer sequences with more memory
|
EPOCHS = 200
|
||||||
PREDICTION_HORIZON = 10 # Number of steps to predict ahead
|
BATCH_SIZE = 128
|
||||||
|
SEQUENCE_LENGTH = 60
|
||||||
|
PREDICTION_HORIZON = 10
|
||||||
|
|
||||||
# Intraday trading settings
|
# Intraday trading settings
|
||||||
TRADING_FREQUENCY = "5min" # '1min', '5min', '15min', '30min', '1h'
|
TRADING_FREQUENCY = "5min" # '1min', '5min', '15min', '30min', '1h'
|
||||||
@@ -127,6 +129,11 @@ class Config:
|
|||||||
NEWS_LOOKBACK_DAYS = 7 # Number of days to look back for news
|
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
|
SOCIAL_MEDIA_LOOKBACK_DAYS = 3 # Number of days to look back for social media
|
||||||
|
|
||||||
|
# Alternative data flag — set to True once NewsProcessor / SocialMediaProcessor
|
||||||
|
# are implemented; False skips those branches so the model trains on real data
|
||||||
|
# rather than zeros, which would mislead the attention layers.
|
||||||
|
USE_ALTERNATIVE_DATA = False
|
||||||
|
|
||||||
# Feature definitions (referenced by models)
|
# Feature definitions (referenced by models)
|
||||||
NEWS_FEATURES = [
|
NEWS_FEATURES = [
|
||||||
"sentiment",
|
"sentiment",
|
||||||
@@ -164,6 +171,14 @@ class Config:
|
|||||||
SLIPPAGE_MODEL = "volume_curve" # 'volume_curve', 'constant', or 'none'
|
SLIPPAGE_MODEL = "volume_curve" # 'volume_curve', 'constant', or 'none'
|
||||||
SLIPPAGE_RATE = 0.0002 # 0.02% slippage
|
SLIPPAGE_RATE = 0.0002 # 0.02% slippage
|
||||||
|
|
||||||
|
# Training loss: weight for (1 − Pearson IC) vs MSE.
|
||||||
|
# 0 = pure MSE, 1 = pure IC loss. 0.7 works well empirically for daily returns.
|
||||||
|
IC_LOSS_WEIGHT = 0.7
|
||||||
|
|
||||||
|
# Walk-forward cross-validation
|
||||||
|
WALK_FORWARD_FOLDS = 5 # number of folds
|
||||||
|
WALK_FORWARD_TEST_YEARS = 1 # out-of-sample window per fold (years)
|
||||||
|
|
||||||
# Live trading settings
|
# Live trading settings
|
||||||
LIVE_DATA_ENABLED = True
|
LIVE_DATA_ENABLED = True
|
||||||
DATA_PROVIDER = "polygon" # 'polygon', 'alphavantage', 'ib', 'tdameritrade'
|
DATA_PROVIDER = "polygon" # 'polygon', 'alphavantage', 'ib', 'tdameritrade'
|
||||||
|
|||||||
+65
-2
@@ -779,7 +779,12 @@ class StockDataPipeline:
|
|||||||
)
|
)
|
||||||
node_tickers = valid_tickers_for_y
|
node_tickers = valid_tickers_for_y
|
||||||
|
|
||||||
y = torch.tensor(y_vals, dtype=torch.float32).unsqueeze(1)
|
y_arr = np.array(y_vals, dtype=np.float32)
|
||||||
|
if len(y_arr) > 1:
|
||||||
|
mu, sigma = y_arr.mean(), y_arr.std()
|
||||||
|
if sigma > 1e-8:
|
||||||
|
y_arr = (y_arr - mu) / sigma
|
||||||
|
y = torch.tensor(y_arr, dtype=torch.float32).unsqueeze(1)
|
||||||
|
|
||||||
data = Data(x=x, edge_index=edge_index, edge_attr=edge_weight, y=y)
|
data = Data(x=x, edge_index=edge_index, edge_attr=edge_weight, y=y)
|
||||||
data.date = date
|
data.date = date
|
||||||
@@ -1001,7 +1006,12 @@ class StockDataPipeline:
|
|||||||
ret = future_price / current_price - 1
|
ret = future_price / current_price - 1
|
||||||
y_vals.append(ret)
|
y_vals.append(ret)
|
||||||
|
|
||||||
y = torch.tensor(y_vals, dtype=torch.float32).unsqueeze(1)
|
y_arr = np.array(y_vals, dtype=np.float32)
|
||||||
|
if len(y_arr) > 1:
|
||||||
|
mu, sigma = y_arr.mean(), y_arr.std()
|
||||||
|
if sigma > 1e-8:
|
||||||
|
y_arr = (y_arr - mu) / sigma
|
||||||
|
y = torch.tensor(y_arr, dtype=torch.float32).unsqueeze(1)
|
||||||
|
|
||||||
# Create Data object
|
# Create Data object
|
||||||
data = Data(x=x, edge_index=edge_index, edge_attr=edge_weight, y=y)
|
data = Data(x=x, edge_index=edge_index, edge_attr=edge_weight, y=y)
|
||||||
@@ -1180,3 +1190,56 @@ class StockDataPipeline:
|
|||||||
"""Create a validation dataset for the given date range."""
|
"""Create a validation dataset for the given date range."""
|
||||||
logger.info(f"Creating validation dataset {start_date} → {end_date}")
|
logger.info(f"Creating validation dataset {start_date} → {end_date}")
|
||||||
return self._create_dataset(start_date, end_date)
|
return self._create_dataset(start_date, end_date)
|
||||||
|
|
||||||
|
def create_walk_forward_splits(self) -> List[Tuple]:
|
||||||
|
"""
|
||||||
|
Build walk-forward CV splits as expanding-window train / fixed OOS val pairs.
|
||||||
|
|
||||||
|
Each fold's training window starts at config.START_DATE and ends at a boundary
|
||||||
|
that advances by config.WALK_FORWARD_TEST_YEARS each fold. The val window is
|
||||||
|
the following config.WALK_FORWARD_TEST_YEARS of data. Folds that would push
|
||||||
|
the val end beyond config.TRAIN_END_DATE (or available data) are dropped.
|
||||||
|
|
||||||
|
Returns a list of (train_dataset, val_dataset, meta_dict) tuples.
|
||||||
|
"""
|
||||||
|
from datetime import datetime, timedelta
|
||||||
|
|
||||||
|
train_start = config.START_DATE
|
||||||
|
full_end = datetime.strptime(config.TRAIN_END_DATE, "%Y-%m-%d")
|
||||||
|
test_years = config.WALK_FORWARD_TEST_YEARS
|
||||||
|
n_folds = config.WALK_FORWARD_FOLDS
|
||||||
|
|
||||||
|
splits: List[Tuple] = []
|
||||||
|
for fold in range(n_folds):
|
||||||
|
# Each fold's val window is offset by fold * test_years from the full_end
|
||||||
|
# minus (n_folds - fold) * test_years so folds are evenly spaced
|
||||||
|
offset_years = (n_folds - fold) * test_years
|
||||||
|
val_end_dt = full_end - timedelta(days=int((offset_years - test_years) * 365))
|
||||||
|
train_end_dt = val_end_dt - timedelta(days=int(test_years * 365))
|
||||||
|
val_start_dt = train_end_dt + timedelta(days=1)
|
||||||
|
|
||||||
|
if train_end_dt <= datetime.strptime(train_start, "%Y-%m-%d"):
|
||||||
|
logger.warning(f"Skipping fold {fold}: train window too short")
|
||||||
|
continue
|
||||||
|
|
||||||
|
train_end_str = train_end_dt.strftime("%Y-%m-%d")
|
||||||
|
val_start_str = val_start_dt.strftime("%Y-%m-%d")
|
||||||
|
val_end_str = val_end_dt.strftime("%Y-%m-%d")
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
f"Walk-forward fold {fold + 1}/{n_folds}: "
|
||||||
|
f"train {train_start}→{train_end_str}, val {val_start_str}→{val_end_str}"
|
||||||
|
)
|
||||||
|
|
||||||
|
train_ds = self._create_dataset(train_start, train_end_str)
|
||||||
|
val_ds = self._create_dataset(val_start_str, val_end_str)
|
||||||
|
meta = {
|
||||||
|
"fold": fold,
|
||||||
|
"train_start": train_start,
|
||||||
|
"train_end": train_end_str,
|
||||||
|
"val_start": val_start_str,
|
||||||
|
"val_end": val_end_str,
|
||||||
|
}
|
||||||
|
splits.append((train_ds, val_ds, meta))
|
||||||
|
|
||||||
|
return splits
|
||||||
|
|||||||
+104
-103
@@ -120,54 +120,67 @@ class TemporalAttention(nn.Module):
|
|||||||
|
|
||||||
class CorporateActionAwareGNN(nn.Module):
|
class CorporateActionAwareGNN(nn.Module):
|
||||||
"""
|
"""
|
||||||
GNN model with corporate action awareness, optimized for AMD Radeon R9700 AI Pro
|
GNN model with corporate action awareness, optimized for AMD Radeon R9700 AI Pro.
|
||||||
|
|
||||||
|
When use_alternative_data=False (the default, driven by config.USE_ALTERNATIVE_DATA)
|
||||||
|
the news and social branches are omitted entirely, which avoids training on the
|
||||||
|
zero-filled stub features returned by the placeholder processors and reduces
|
||||||
|
the parameter count to match the data actually available.
|
||||||
|
Set to True once real NewsProcessor / SocialMediaProcessor implementations exist.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, num_node_features: int):
|
def __init__(self, num_node_features: int, use_alternative_data: bool = None):
|
||||||
super(CorporateActionAwareGNN, self).__init__()
|
super(CorporateActionAwareGNN, self).__init__()
|
||||||
|
|
||||||
# Calculate feature dimensions
|
self.use_alternative_data = (
|
||||||
num_price_features = 5 # returns, volatility, momentum, volume, price
|
config.USE_ALTERNATIVE_DATA
|
||||||
num_news_features = len(config.NEWS_FEATURES)
|
if use_alternative_data is None
|
||||||
num_social_features = len(config.SOCIAL_FEATURES)
|
else use_alternative_data
|
||||||
|
)
|
||||||
|
|
||||||
# Feature processing modules with AMD optimizations
|
H = config.HIDDEN_CHANNELS
|
||||||
|
|
||||||
|
# Price and corporate-action processors — always active
|
||||||
self.price_processor = nn.Sequential(
|
self.price_processor = nn.Sequential(
|
||||||
nn.Linear(num_price_features, config.HIDDEN_CHANNELS),
|
nn.Linear(5, H),
|
||||||
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.SiLU(),
|
||||||
nn.Linear(config.HIDDEN_CHANNELS, config.HIDDEN_CHANNELS),
|
nn.Linear(H, H),
|
||||||
nn.LayerNorm(config.HIDDEN_CHANNELS),
|
nn.LayerNorm(H),
|
||||||
)
|
)
|
||||||
|
|
||||||
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(
|
self.corporate_action_mlp = nn.Sequential(
|
||||||
nn.Linear(1, config.HIDDEN_CHANNELS),
|
nn.Linear(1, H),
|
||||||
nn.SiLU(),
|
nn.SiLU(),
|
||||||
nn.Linear(config.HIDDEN_CHANNELS, config.HIDDEN_CHANNELS),
|
nn.Linear(H, H),
|
||||||
nn.LayerNorm(config.HIDDEN_CHANNELS),
|
nn.LayerNorm(H),
|
||||||
)
|
)
|
||||||
|
self.temporal_attention = TemporalAttention(H, num_heads=8)
|
||||||
|
|
||||||
# Temporal attention for sequence processing
|
# Optional alternative-data branches (news + social)
|
||||||
self.temporal_attention = TemporalAttention(config.HIDDEN_CHANNELS, num_heads=8)
|
if self.use_alternative_data:
|
||||||
|
num_news = len(config.NEWS_FEATURES)
|
||||||
|
num_social = len(config.SOCIAL_FEATURES)
|
||||||
|
self.news_processor = nn.Sequential(
|
||||||
|
nn.Linear(num_news, H), nn.SiLU(), nn.Linear(H, H), nn.LayerNorm(H)
|
||||||
|
)
|
||||||
|
self.social_processor = nn.Sequential(
|
||||||
|
nn.Linear(num_social, H), nn.SiLU(), nn.Linear(H, H), nn.LayerNorm(H)
|
||||||
|
)
|
||||||
|
# 3-way softmax over (price, news, social)
|
||||||
|
self.alternative_data_attention = nn.Sequential(
|
||||||
|
nn.Linear(H * 3, H),
|
||||||
|
nn.SiLU(),
|
||||||
|
nn.Linear(H, 3),
|
||||||
|
nn.Softmax(dim=1),
|
||||||
|
)
|
||||||
|
gnn_in = H * 4 # price + news + social + corporate
|
||||||
|
else:
|
||||||
|
gnn_in = H * 2 # price + corporate only
|
||||||
|
|
||||||
# Graph attention layers with AMD optimizations
|
# Graph attention layers with AMD optimizations
|
||||||
try:
|
try:
|
||||||
self.conv1 = AMDGATConv(
|
self.conv1 = AMDGATConv(
|
||||||
config.HIDDEN_CHANNELS * 4, # Combined features from all processors
|
gnn_in,
|
||||||
config.HIDDEN_CHANNELS,
|
H,
|
||||||
heads=config.NUM_HEADS,
|
heads=config.NUM_HEADS,
|
||||||
concat=True,
|
concat=True,
|
||||||
dropout=config.DROPOUT,
|
dropout=config.DROPOUT,
|
||||||
@@ -175,8 +188,8 @@ class CorporateActionAwareGNN(nn.Module):
|
|||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
self.conv1 = GATConv(
|
self.conv1 = GATConv(
|
||||||
config.HIDDEN_CHANNELS * 4,
|
gnn_in,
|
||||||
config.HIDDEN_CHANNELS,
|
H,
|
||||||
heads=config.NUM_HEADS,
|
heads=config.NUM_HEADS,
|
||||||
concat=True,
|
concat=True,
|
||||||
dropout=config.DROPOUT,
|
dropout=config.DROPOUT,
|
||||||
@@ -185,8 +198,8 @@ class CorporateActionAwareGNN(nn.Module):
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
self.conv2 = AMDGATConv(
|
self.conv2 = AMDGATConv(
|
||||||
config.HIDDEN_CHANNELS * config.NUM_HEADS,
|
H * config.NUM_HEADS,
|
||||||
config.HIDDEN_CHANNELS,
|
H,
|
||||||
heads=1,
|
heads=1,
|
||||||
concat=False,
|
concat=False,
|
||||||
dropout=config.DROPOUT,
|
dropout=config.DROPOUT,
|
||||||
@@ -194,44 +207,32 @@ class CorporateActionAwareGNN(nn.Module):
|
|||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
self.conv2 = GATConv(
|
self.conv2 = GATConv(
|
||||||
config.HIDDEN_CHANNELS * config.NUM_HEADS,
|
H * config.NUM_HEADS,
|
||||||
config.HIDDEN_CHANNELS,
|
H,
|
||||||
heads=1,
|
heads=1,
|
||||||
concat=False,
|
concat=False,
|
||||||
dropout=config.DROPOUT,
|
dropout=config.DROPOUT,
|
||||||
add_self_loops=True,
|
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),
|
|
||||||
)
|
|
||||||
|
|
||||||
# MLP applied after GNN (replaces single-timestep LSTM which was equivalent
|
|
||||||
# to a linear layer but 4× more parameters)
|
|
||||||
self.post_gnn_mlp = nn.Sequential(
|
self.post_gnn_mlp = nn.Sequential(
|
||||||
nn.Linear(config.HIDDEN_CHANNELS, config.HIDDEN_CHANNELS),
|
nn.Linear(H, H),
|
||||||
nn.SiLU(),
|
nn.SiLU(),
|
||||||
nn.LayerNorm(config.HIDDEN_CHANNELS),
|
nn.LayerNorm(H),
|
||||||
)
|
)
|
||||||
|
# GNN output (H) concat with modal summary (H) → 1
|
||||||
|
self.linear = nn.Linear(H * 2, 1)
|
||||||
|
|
||||||
# Final prediction: GNN output (HIDDEN) concat with weighted modal sum (HIDDEN)
|
|
||||||
self.linear = nn.Linear(config.HIDDEN_CHANNELS * 2, 1)
|
|
||||||
|
|
||||||
# Initialize weights
|
|
||||||
self._init_weights()
|
self._init_weights()
|
||||||
|
|
||||||
def _init_weights(self):
|
def _init_weights(self):
|
||||||
"""Initialize weights with AMD-friendly initialization"""
|
"""Initialize weights with AMD-friendly initialization."""
|
||||||
nn.init.xavier_uniform_(self.linear.weight)
|
nn.init.xavier_uniform_(self.linear.weight)
|
||||||
nn.init.zeros_(self.linear.bias)
|
nn.init.zeros_(self.linear.bias)
|
||||||
|
modules_to_init = [self.post_gnn_mlp]
|
||||||
for module in (self.alternative_data_attention, self.post_gnn_mlp):
|
if self.use_alternative_data:
|
||||||
|
modules_to_init.append(self.alternative_data_attention)
|
||||||
|
for module in modules_to_init:
|
||||||
for layer in module:
|
for layer in module:
|
||||||
if isinstance(layer, nn.Linear):
|
if isinstance(layer, nn.Linear):
|
||||||
nn.init.xavier_uniform_(layer.weight)
|
nn.init.xavier_uniform_(layer.weight)
|
||||||
@@ -251,72 +252,72 @@ class CorporateActionAwareGNN(nn.Module):
|
|||||||
|
|
||||||
def forward(self, data):
|
def forward(self, data):
|
||||||
x, edge_index, edge_attr = data.x, data.edge_index, data.edge_attr
|
x, edge_index, edge_attr = data.x, data.edge_index, data.edge_attr
|
||||||
|
|
||||||
price_features, news_features, social_features, corporate_action_flags = (
|
price_features, news_features, social_features, corporate_action_flags = (
|
||||||
self._split_features(x)
|
self._split_features(x)
|
||||||
)
|
)
|
||||||
|
|
||||||
price_processed = self.price_processor(price_features)
|
price_attended = self.temporal_attention(self.price_processor(price_features))
|
||||||
news_processed = self.news_processor(news_features)
|
corporate_attended = self.temporal_attention(
|
||||||
social_processed = self.social_processor(social_features)
|
self.corporate_action_mlp(corporate_action_flags)
|
||||||
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)
|
|
||||||
corporate_attended = self.temporal_attention(corporate_action_features)
|
|
||||||
|
|
||||||
alternative_features = torch.cat(
|
|
||||||
[price_attended, news_attended, social_attended], dim=1
|
|
||||||
)
|
)
|
||||||
attention_weights = self.alternative_data_attention(alternative_features)
|
|
||||||
|
|
||||||
weighted_price = price_attended * attention_weights[:, 0].unsqueeze(1)
|
if self.use_alternative_data:
|
||||||
weighted_news = news_attended * attention_weights[:, 1].unsqueeze(1)
|
news_attended = self.temporal_attention(
|
||||||
weighted_social = social_attended * attention_weights[:, 2].unsqueeze(1)
|
self.news_processor(news_features)
|
||||||
|
)
|
||||||
combined_features = torch.cat(
|
social_attended = self.temporal_attention(
|
||||||
[weighted_price, weighted_news, weighted_social, corporate_attended], dim=1
|
self.social_processor(social_features)
|
||||||
)
|
)
|
||||||
|
alt_feats = torch.cat(
|
||||||
|
[price_attended, news_attended, social_attended], dim=1
|
||||||
|
)
|
||||||
|
attn_weights = self.alternative_data_attention(alt_feats)
|
||||||
|
weighted_price = price_attended * attn_weights[:, 0].unsqueeze(1)
|
||||||
|
weighted_news = news_attended * attn_weights[:, 1].unsqueeze(1)
|
||||||
|
weighted_social = social_attended * attn_weights[:, 2].unsqueeze(1)
|
||||||
|
combined_features = torch.cat(
|
||||||
|
[weighted_price, weighted_news, weighted_social, corporate_attended],
|
||||||
|
dim=1,
|
||||||
|
)
|
||||||
|
modal_combined = weighted_price + weighted_news + weighted_social
|
||||||
|
else:
|
||||||
|
combined_features = torch.cat(
|
||||||
|
[price_attended, corporate_attended], dim=1
|
||||||
|
)
|
||||||
|
modal_combined = price_attended
|
||||||
|
|
||||||
x = self._gnn_forward(combined_features, edge_index, edge_attr)
|
x = self._gnn_forward(combined_features, edge_index, edge_attr)
|
||||||
x = self.post_gnn_mlp(x) # (num_stocks, HIDDEN_CHANNELS)
|
x = self.post_gnn_mlp(x)
|
||||||
|
|
||||||
# Weighted combination of attended modalities: (N, HIDDEN_CHANNELS)
|
|
||||||
modal_combined = weighted_price + weighted_news + weighted_social
|
|
||||||
|
|
||||||
# Concat GNN output with modal summary → (N, HIDDEN_CHANNELS * 2)
|
|
||||||
x = torch.cat([x, modal_combined], dim=1)
|
x = torch.cat([x, modal_combined], dim=1)
|
||||||
|
|
||||||
return self.linear(x)
|
return self.linear(x)
|
||||||
|
|
||||||
def _gnn_forward(self, x, edge_index, edge_attr):
|
def _gnn_forward(self, x, edge_index, edge_attr):
|
||||||
"""Forward pass through GNN with gradient checkpointing for memory efficiency"""
|
"""Forward through GNN with gradient checkpointing for memory efficiency."""
|
||||||
# Gradient checkpointing for memory efficiency on AMD GPUs
|
|
||||||
x = torch.utils.checkpoint.checkpoint(
|
x = torch.utils.checkpoint.checkpoint(
|
||||||
self.conv1, x, edge_index, edge_attr, preserve_rng_state=False
|
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.silu(x)
|
||||||
x = F.dropout(x, p=config.DROPOUT, training=self.training)
|
x = F.dropout(x, p=config.DROPOUT, training=self.training)
|
||||||
|
|
||||||
x = torch.utils.checkpoint.checkpoint(
|
x = torch.utils.checkpoint.checkpoint(
|
||||||
self.conv2, x, edge_index, edge_attr, preserve_rng_state=False
|
self.conv2, x, edge_index, edge_attr, preserve_rng_state=False
|
||||||
)
|
)
|
||||||
|
|
||||||
return x
|
return x
|
||||||
|
|
||||||
def get_attention_weights(self, data):
|
def get_attention_weights(self, data):
|
||||||
"""Get attention weights for interpretability"""
|
"""Get modality attention weights for interpretability."""
|
||||||
x = data.x
|
x = data.x
|
||||||
price_features, news_features, social_features, _ = self._split_features(x)
|
price_features, news_features, social_features, _ = self._split_features(x)
|
||||||
|
|
||||||
price_attended = self.temporal_attention(self.price_processor(price_features))
|
price_attended = self.temporal_attention(self.price_processor(price_features))
|
||||||
news_attended = self.temporal_attention(self.news_processor(news_features))
|
if self.use_alternative_data:
|
||||||
social_attended = self.temporal_attention(
|
news_attended = self.temporal_attention(
|
||||||
self.social_processor(social_features)
|
self.news_processor(news_features)
|
||||||
)
|
)
|
||||||
|
social_attended = self.temporal_attention(
|
||||||
alternative_features = torch.cat(
|
self.social_processor(social_features)
|
||||||
[price_attended, news_attended, social_attended], dim=1
|
)
|
||||||
)
|
alt_feats = torch.cat(
|
||||||
return self.alternative_data_attention(alternative_features)
|
[price_attended, news_attended, social_attended], dim=1
|
||||||
|
)
|
||||||
|
return self.alternative_data_attention(alt_feats)
|
||||||
|
# Price-only: return uniform weight of 1.0 per stock
|
||||||
|
return price_attended.new_ones(price_attended.size(0), 1)
|
||||||
|
|||||||
+70
-3
@@ -6,6 +6,7 @@ from typing import Dict, List, Tuple
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
from torch.amp import autocast
|
from torch.amp import autocast
|
||||||
from torch.cuda.amp import GradScaler
|
from torch.cuda.amp import GradScaler
|
||||||
from torch_geometric.loader import DataLoader
|
from torch_geometric.loader import DataLoader
|
||||||
@@ -35,7 +36,7 @@ class GNNTrainer:
|
|||||||
|
|
||||||
# Set up learning rate scheduler
|
# Set up learning rate scheduler
|
||||||
self.scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
|
self.scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
|
||||||
self.optimizer, mode="min", factor=0.5, patience=5, verbose=True
|
self.optimizer, mode="min", factor=0.5, patience=5
|
||||||
)
|
)
|
||||||
|
|
||||||
self.criterion = nn.MSELoss()
|
self.criterion = nn.MSELoss()
|
||||||
@@ -117,7 +118,7 @@ class GNNTrainer:
|
|||||||
dtype=self.amd_optimizer.get_precision_dtype(),
|
dtype=self.amd_optimizer.get_precision_dtype(),
|
||||||
):
|
):
|
||||||
out = self.model(batch)
|
out = self.model(batch)
|
||||||
loss = self.criterion(out, batch.y)
|
loss = self._compute_loss(out, batch.y)
|
||||||
|
|
||||||
# Scale loss and backpropagate
|
# Scale loss and backpropagate
|
||||||
self.scaler.scale(loss).backward()
|
self.scaler.scale(loss).backward()
|
||||||
@@ -171,7 +172,7 @@ class GNNTrainer:
|
|||||||
dtype=self.amd_optimizer.get_precision_dtype(),
|
dtype=self.amd_optimizer.get_precision_dtype(),
|
||||||
):
|
):
|
||||||
out = self.model(batch)
|
out = self.model(batch)
|
||||||
loss = self.criterion(out, batch.y)
|
loss = self._compute_loss(out, batch.y)
|
||||||
|
|
||||||
epoch_val_loss += loss.item()
|
epoch_val_loss += loss.item()
|
||||||
num_val_batches += 1
|
num_val_batches += 1
|
||||||
@@ -364,6 +365,72 @@ class GNNTrainer:
|
|||||||
self.memory_manager.empty_cache()
|
self.memory_manager.empty_cache()
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
def _ic_loss(self, pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""1 − Pearson IC: differentiable surrogate for rank correlation."""
|
||||||
|
p = pred.flatten()
|
||||||
|
t = target.flatten()
|
||||||
|
if p.numel() < 2:
|
||||||
|
return F.mse_loss(pred, target)
|
||||||
|
p_c = p - p.mean()
|
||||||
|
t_c = t - t.mean()
|
||||||
|
ic = (p_c * t_c).sum() / (p_c.norm() * t_c.norm() + 1e-8)
|
||||||
|
return 1.0 - ic
|
||||||
|
|
||||||
|
def _compute_loss(self, pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
|
||||||
|
mse = self.criterion(pred, target)
|
||||||
|
if config.IC_LOSS_WEIGHT <= 0:
|
||||||
|
return mse
|
||||||
|
return config.IC_LOSS_WEIGHT * self._ic_loss(pred, target) + (1.0 - config.IC_LOSS_WEIGHT) * mse
|
||||||
|
|
||||||
|
def _reset_model_weights(self):
|
||||||
|
"""Re-initialise all model parameters (for walk-forward resets)."""
|
||||||
|
for module in self.model.modules():
|
||||||
|
if hasattr(module, "reset_parameters"):
|
||||||
|
module.reset_parameters()
|
||||||
|
|
||||||
|
def walk_forward_train(self, splits: List[Tuple], reset_weights: bool = True) -> Dict:
|
||||||
|
"""
|
||||||
|
Train on each walk-forward fold and return per-fold validation losses.
|
||||||
|
|
||||||
|
splits: list of (train_dataset, val_dataset, metadata) tuples as returned by
|
||||||
|
DataPipeline.create_walk_forward_splits()
|
||||||
|
"""
|
||||||
|
fold_results = {}
|
||||||
|
original_best = self.best_val_loss
|
||||||
|
|
||||||
|
for fold_idx, (train_ds, val_ds, meta) in enumerate(splits):
|
||||||
|
logger.info(
|
||||||
|
f"Walk-forward fold {fold_idx + 1}/{len(splits)} — "
|
||||||
|
f"train end: {meta.get('train_end')}, val: {meta.get('val_start')}–{meta.get('val_end')}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if reset_weights:
|
||||||
|
self._reset_model_weights()
|
||||||
|
# Re-init optimizer so momentum buffers don't bleed across folds
|
||||||
|
self.optimizer = torch.optim.AdamW(
|
||||||
|
self.model.parameters(), lr=config.LEARNING_RATE, weight_decay=1e-4
|
||||||
|
)
|
||||||
|
self.scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
|
||||||
|
self.optimizer, mode="min", factor=0.5, patience=5
|
||||||
|
)
|
||||||
|
self.scaler = GradScaler(enabled=config.MIXED_PRECISION)
|
||||||
|
|
||||||
|
self.train_losses = []
|
||||||
|
self.val_losses = []
|
||||||
|
self.best_val_loss = float("inf")
|
||||||
|
|
||||||
|
train_losses, val_losses = self.train(train_ds, val_ds)
|
||||||
|
fold_results[fold_idx] = {
|
||||||
|
"train_losses": train_losses,
|
||||||
|
"val_losses": val_losses,
|
||||||
|
"best_val_loss": self.best_val_loss,
|
||||||
|
"meta": meta,
|
||||||
|
}
|
||||||
|
logger.info(f"Fold {fold_idx + 1} best val loss: {self.best_val_loss:.6f}")
|
||||||
|
|
||||||
|
self.best_val_loss = original_best
|
||||||
|
return fold_results
|
||||||
|
|
||||||
def benchmark(self, input_data, num_runs: int = 100):
|
def benchmark(self, input_data, num_runs: int = 100):
|
||||||
"""Benchmark model performance on AMD GPU"""
|
"""Benchmark model performance on AMD GPU"""
|
||||||
return self.amd_optimizer.benchmark_model(self.model, input_data, num_runs)
|
return self.amd_optimizer.benchmark_model(self.model, input_data, num_runs)
|
||||||
|
|||||||
Reference in New Issue
Block a user