460 lines
15 KiB
Python
460 lines
15 KiB
Python
#!/usr/bin/env python3
|
|
"""
|
|
gnn_model.py - Graph Neural Network for Cross-Asset Trading
|
|
=============================================================
|
|
Models the crypto market as a graph where:
|
|
- Each coin is a NODE with its features
|
|
- EDGES represent correlations between coins
|
|
- Message passing captures lead-lag relationships
|
|
|
|
Architecture:
|
|
1. Graph construction from correlation matrix
|
|
2. GCN (Graph Convolutional Network) layers for message passing
|
|
3. Node-level prediction (per-coin UP/DOWN probability)
|
|
|
|
Key insight: if BTC drops, the GNN learns that ETH follows in ~5min,
|
|
SOL in ~10min, DOGE in ~30min. This gives earlier signals.
|
|
|
|
Dependencies: torch (PyTorch). No torch_geometric needed.
|
|
"""
|
|
import numpy as np
|
|
import pickle
|
|
import logging
|
|
|
|
log = logging.getLogger("AHAD QUANT")
|
|
|
|
try:
|
|
import torch
|
|
import torch.nn as nn
|
|
import torch.nn.functional as F
|
|
HAS_TORCH = True
|
|
except ImportError:
|
|
HAS_TORCH = False
|
|
|
|
|
|
class GraphConvLayer(nn.Module):
|
|
"""Simple Graph Convolutional Layer (Kipf & Welling, 2017)."""
|
|
def __init__(self, in_features, out_features):
|
|
super().__init__()
|
|
self.weight = nn.Parameter(torch.FloatTensor(in_features, out_features))
|
|
self.bias = nn.Parameter(torch.FloatTensor(out_features))
|
|
nn.init.xavier_uniform_(self.weight)
|
|
nn.init.zeros_(self.bias)
|
|
|
|
def forward(self, x, adj):
|
|
"""
|
|
Args:
|
|
x: (batch, n_nodes, in_features)
|
|
adj: (n_nodes, n_nodes) normalized adjacency matrix
|
|
|
|
Returns:
|
|
(batch, n_nodes, out_features)
|
|
"""
|
|
# Message passing: A * X * W + b
|
|
support = torch.matmul(x, self.weight) # (batch, n_nodes, out_features)
|
|
output = torch.matmul(adj, support) # (batch, n_nodes, out_features)
|
|
return output + self.bias
|
|
|
|
|
|
class CryptoGNN(nn.Module):
|
|
"""
|
|
Graph Neural Network for crypto market.
|
|
|
|
Takes features for ALL coins simultaneously and predicts
|
|
UP/DOWN probability for each coin, using cross-asset information.
|
|
"""
|
|
def __init__(self, n_features, hidden_dim=64, n_gcn_layers=3, dropout=0.2):
|
|
super().__init__()
|
|
self.n_features = n_features
|
|
|
|
# Node feature encoder (shared across all coins)
|
|
self.encoder = nn.Sequential(
|
|
nn.Linear(n_features, hidden_dim),
|
|
nn.ReLU(),
|
|
nn.Dropout(dropout),
|
|
)
|
|
|
|
# GCN layers
|
|
self.gcn_layers = nn.ModuleList()
|
|
self.gcn_norms = nn.ModuleList()
|
|
for i in range(n_gcn_layers):
|
|
self.gcn_layers.append(GraphConvLayer(hidden_dim, hidden_dim))
|
|
self.gcn_norms.append(nn.LayerNorm(hidden_dim))
|
|
|
|
self.dropout = nn.Dropout(dropout)
|
|
|
|
# Node-level prediction head
|
|
self.head = nn.Sequential(
|
|
nn.Linear(hidden_dim * 2, hidden_dim), # concat local + global
|
|
nn.ReLU(),
|
|
nn.Dropout(dropout),
|
|
nn.Linear(hidden_dim, 1),
|
|
nn.Sigmoid()
|
|
)
|
|
|
|
def forward(self, x, adj):
|
|
"""
|
|
Args:
|
|
x: (batch, n_nodes, n_features) - features for all coins
|
|
adj: (n_nodes, n_nodes) - normalized adjacency matrix
|
|
|
|
Returns:
|
|
pred: (batch, n_nodes, 1) - UP probability per coin
|
|
"""
|
|
# Encode node features
|
|
h = self.encoder(x) # (batch, n_nodes, hidden)
|
|
|
|
# GCN message passing
|
|
for gcn, norm in zip(self.gcn_layers, self.gcn_norms):
|
|
h_new = gcn(h, adj)
|
|
h_new = F.relu(h_new)
|
|
h_new = self.dropout(h_new)
|
|
h = norm(h_new + h) # residual connection
|
|
|
|
# Global graph context (mean pooling)
|
|
global_ctx = h.mean(dim=1, keepdim=True).expand_as(h) # (batch, n_nodes, hidden)
|
|
|
|
# Concatenate local + global for prediction
|
|
combined = torch.cat([h, global_ctx], dim=-1) # (batch, n_nodes, hidden*2)
|
|
pred = self.head(combined) # (batch, n_nodes, 1)
|
|
|
|
return pred
|
|
|
|
|
|
def build_correlation_graph(returns_dict, threshold=0.3):
|
|
"""
|
|
Build adjacency matrix from return correlations.
|
|
|
|
Args:
|
|
returns_dict: dict of {coin_name: np.array of returns}
|
|
threshold: minimum absolute correlation for edge (default 0.3)
|
|
|
|
Returns:
|
|
adj: (n_coins, n_coins) normalized adjacency matrix
|
|
coin_order: list of coin names in matrix order
|
|
"""
|
|
coins = sorted(returns_dict.keys())
|
|
n = len(coins)
|
|
|
|
# Compute correlation matrix
|
|
returns_matrix = np.column_stack([returns_dict[c] for c in coins])
|
|
min_len = min(len(returns_dict[c]) for c in coins)
|
|
returns_matrix = returns_matrix[:min_len]
|
|
|
|
corr = np.corrcoef(returns_matrix.T)
|
|
corr = np.nan_to_num(corr, nan=0.0)
|
|
|
|
# Build adjacency: threshold + self-loops
|
|
adj = np.zeros((n, n))
|
|
for i in range(n):
|
|
adj[i, i] = 1.0 # self-loop
|
|
for j in range(i + 1, n):
|
|
if abs(corr[i, j]) > threshold:
|
|
adj[i, j] = abs(corr[i, j])
|
|
adj[j, i] = abs(corr[i, j])
|
|
|
|
# Normalize: D^{-1/2} A D^{-1/2}
|
|
degree = adj.sum(axis=1)
|
|
d_inv_sqrt = np.zeros_like(degree)
|
|
nonzero = degree > 0
|
|
d_inv_sqrt[nonzero] = 1.0 / np.sqrt(degree[nonzero])
|
|
D = np.diag(d_inv_sqrt)
|
|
adj_norm = D @ adj @ D
|
|
|
|
return adj_norm.astype(np.float32), coins
|
|
|
|
|
|
def build_lead_lag_graph(returns_dict, max_lag=5):
|
|
"""
|
|
Build directed graph based on lead-lag relationships.
|
|
|
|
If coin A's returns at time t predict coin B at time t+lag,
|
|
then A→B edge exists.
|
|
|
|
Args:
|
|
returns_dict: dict of {coin: returns_array}
|
|
max_lag: maximum lag to check (in candles)
|
|
|
|
Returns:
|
|
adj: (n_coins, n_coins) normalized adjacency (DIRECTED)
|
|
coin_order: list of coin names
|
|
"""
|
|
coins = sorted(returns_dict.keys())
|
|
n = len(coins)
|
|
min_len = min(len(returns_dict[c]) for c in coins)
|
|
|
|
adj = np.eye(n, dtype=np.float32) # self-loops
|
|
|
|
for i, ci in enumerate(coins):
|
|
ri = returns_dict[ci][:min_len]
|
|
for j, cj in enumerate(coins):
|
|
if i == j:
|
|
continue
|
|
rj = returns_dict[cj][:min_len]
|
|
|
|
# Check if coin i leads coin j
|
|
best_corr = 0.0
|
|
for lag in range(1, max_lag + 1):
|
|
if lag >= min_len:
|
|
break
|
|
corr = np.corrcoef(ri[:-lag], rj[lag:])[0, 1]
|
|
if not np.isnan(corr):
|
|
best_corr = max(best_corr, abs(corr))
|
|
|
|
if best_corr > 0.15:
|
|
adj[i, j] = best_corr
|
|
|
|
# Row-normalize
|
|
row_sums = adj.sum(axis=1, keepdims=True)
|
|
row_sums = np.where(row_sums > 0, row_sums, 1.0)
|
|
adj_norm = adj / row_sums
|
|
|
|
return adj_norm, coins
|
|
|
|
|
|
def train_gnn(features_dict, labels_dict, adj, coin_order,
|
|
n_features=None, hidden_dim=64, epochs=100,
|
|
lr=0.001, batch_size=32, patience=15, device='cpu'):
|
|
"""
|
|
Train GNN model.
|
|
|
|
Args:
|
|
features_dict: {coin: (n_samples, n_features)} aligned by time
|
|
labels_dict: {coin: (n_samples,)} binary labels
|
|
adj: (n_coins, n_coins) adjacency matrix
|
|
coin_order: list of coin names matching adj rows
|
|
|
|
Returns:
|
|
model: trained CryptoGNN
|
|
history: training metrics
|
|
"""
|
|
if not HAS_TORCH:
|
|
raise ImportError("PyTorch required")
|
|
|
|
# Build aligned data matrix: (n_times, n_coins, n_features)
|
|
n_coins = len(coin_order)
|
|
min_len = min(len(features_dict[c]) for c in coin_order)
|
|
if n_features is None:
|
|
n_features = features_dict[coin_order[0]].shape[1]
|
|
|
|
X_all = np.zeros((min_len, n_coins, n_features), dtype=np.float32)
|
|
y_all = np.zeros((min_len, n_coins), dtype=np.float32)
|
|
|
|
for idx, coin in enumerate(coin_order):
|
|
X_all[:, idx, :] = features_dict[coin][:min_len, :n_features]
|
|
y_all[:, idx] = labels_dict[coin][:min_len]
|
|
|
|
# Split train/val (time-based)
|
|
split = int(min_len * 0.8)
|
|
X_train = X_all[:split]
|
|
y_train = y_all[:split]
|
|
X_val = X_all[split:]
|
|
y_val = y_all[split:]
|
|
|
|
model = CryptoGNN(n_features=n_features, hidden_dim=hidden_dim).to(device)
|
|
adj_tensor = torch.FloatTensor(adj).to(device)
|
|
optimizer = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=1e-5)
|
|
criterion = nn.BCELoss()
|
|
|
|
best_val_loss = float('inf')
|
|
best_state = None
|
|
patience_counter = 0
|
|
history = {'train_loss': [], 'val_loss': [], 'val_acc': []}
|
|
|
|
n_train = len(X_train)
|
|
|
|
for epoch in range(epochs):
|
|
model.train()
|
|
train_loss = 0.0
|
|
n_batches = 0
|
|
|
|
indices = np.random.permutation(n_train)
|
|
for start in range(0, n_train, batch_size):
|
|
end = min(start + batch_size, n_train)
|
|
batch_idx = indices[start:end]
|
|
|
|
x_b = torch.FloatTensor(X_train[batch_idx]).to(device)
|
|
y_b = torch.FloatTensor(y_train[batch_idx]).unsqueeze(-1).to(device)
|
|
|
|
optimizer.zero_grad()
|
|
pred = model(x_b, adj_tensor)
|
|
loss = criterion(pred, y_b)
|
|
loss.backward()
|
|
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
|
|
optimizer.step()
|
|
|
|
train_loss += loss.item()
|
|
n_batches += 1
|
|
|
|
avg_train = train_loss / max(n_batches, 1)
|
|
|
|
# Validation
|
|
model.eval()
|
|
with torch.no_grad():
|
|
x_v = torch.FloatTensor(X_val).to(device)
|
|
y_v = torch.FloatTensor(y_val).unsqueeze(-1).to(device)
|
|
pred_v = model(x_v, adj_tensor)
|
|
val_loss = criterion(pred_v, y_v).item()
|
|
val_acc = ((pred_v > 0.5).float() == y_v).float().mean().item() * 100
|
|
|
|
history['train_loss'].append(avg_train)
|
|
history['val_loss'].append(val_loss)
|
|
history['val_acc'].append(val_acc)
|
|
|
|
if val_loss < best_val_loss:
|
|
best_val_loss = val_loss
|
|
best_state = {k: v.cpu().clone() for k, v in model.state_dict().items()}
|
|
patience_counter = 0
|
|
else:
|
|
patience_counter += 1
|
|
|
|
if (epoch + 1) % 10 == 0:
|
|
log.info(f"[GNN] Epoch {epoch+1}/{epochs} | train={avg_train:.4f} | "
|
|
f"val={val_loss:.4f} | acc={val_acc:.1f}%")
|
|
|
|
if patience_counter >= patience:
|
|
log.info(f"[GNN] Early stopping at epoch {epoch+1}")
|
|
break
|
|
|
|
if best_state:
|
|
model.load_state_dict(best_state)
|
|
model.eval()
|
|
|
|
return model, history
|
|
|
|
|
|
def save_gnn(model, adj, coin_order, config, normalization, accuracy, path):
|
|
"""Save GNN model + graph structure."""
|
|
data = {
|
|
'model_state': model.state_dict(),
|
|
'adj': adj,
|
|
'coin_order': coin_order,
|
|
'config': config,
|
|
'normalization': normalization,
|
|
'accuracy': accuracy,
|
|
'model_type': 'gnn',
|
|
}
|
|
with open(path, 'wb') as f:
|
|
pickle.dump(data, f)
|
|
|
|
|
|
def load_gnn(path):
|
|
"""Load GNN model."""
|
|
if not HAS_TORCH:
|
|
return None, None, None, None
|
|
try:
|
|
with open(path, 'rb') as f:
|
|
data = pickle.load(f)
|
|
cfg = data['config']
|
|
model = CryptoGNN(
|
|
n_features=cfg['n_features'],
|
|
hidden_dim=cfg.get('hidden_dim', 64),
|
|
)
|
|
model.load_state_dict(data['model_state'])
|
|
model.eval()
|
|
return model, data['adj'], data['coin_order'], data.get('normalization', {})
|
|
except Exception as e:
|
|
log.warning(f"[GNN] Load failed: {e}")
|
|
return None, None, None, None
|
|
|
|
|
|
def predict_gnn(model, features_dict, adj, coin_order, normalization=None, target_coin=None):
|
|
"""
|
|
Predict UP probability for all coins (or a specific coin).
|
|
|
|
Args:
|
|
model: trained CryptoGNN
|
|
features_dict: {coin: (n_features,) array} current features for each coin
|
|
adj: adjacency matrix
|
|
coin_order: list of coin names
|
|
normalization: optional {'mean': array, 'std': array}
|
|
target_coin: optional specific coin to get prediction for
|
|
|
|
Returns:
|
|
dict of {coin: probability} or single float if target_coin specified
|
|
"""
|
|
if not HAS_TORCH or model is None:
|
|
return None
|
|
|
|
try:
|
|
n_coins = len(coin_order)
|
|
n_features = model.n_features
|
|
|
|
# Build input tensor
|
|
x = np.zeros((1, n_coins, n_features), dtype=np.float32)
|
|
for idx, coin in enumerate(coin_order):
|
|
if coin in features_dict:
|
|
feat = np.array(features_dict[coin][:n_features], dtype=np.float32)
|
|
if normalization:
|
|
mean = np.array(normalization['mean'][:n_features], dtype=np.float32)
|
|
std = np.array(normalization['std'][:n_features], dtype=np.float32)
|
|
std = np.where(std < 1e-8, 1.0, std)
|
|
feat = (feat - mean) / std
|
|
x[0, idx, :len(feat)] = feat
|
|
|
|
adj_tensor = torch.FloatTensor(adj)
|
|
x_tensor = torch.FloatTensor(x)
|
|
|
|
with torch.no_grad():
|
|
pred = model(x_tensor, adj_tensor) # (1, n_coins, 1)
|
|
|
|
probs = pred[0, :, 0].numpy()
|
|
result = {coin: float(probs[idx]) for idx, coin in enumerate(coin_order)}
|
|
|
|
if target_coin:
|
|
return result.get(target_coin, 0.5)
|
|
return result
|
|
|
|
except Exception as e:
|
|
log.warning(f"[GNN] Predict error: {e}")
|
|
return None
|
|
|
|
|
|
if __name__ == '__main__':
|
|
if not HAS_TORCH:
|
|
print("[GNN] PyTorch not available")
|
|
else:
|
|
print("[GNN] Graph Neural Network smoke test")
|
|
|
|
np.random.seed(42)
|
|
coins = ['BTC', 'ETH', 'SOL', 'DOGE', 'LINK']
|
|
n = 500
|
|
n_features = 20
|
|
|
|
# Synthetic correlated returns
|
|
btc_returns = np.random.randn(n) * 0.01
|
|
returns_dict = {'BTC': btc_returns}
|
|
for coin in coins[1:]:
|
|
lag = np.random.randint(1, 4)
|
|
noise = np.random.randn(n) * 0.005
|
|
r = np.roll(btc_returns, lag) * (0.5 + np.random.rand() * 0.5) + noise
|
|
returns_dict[coin] = r
|
|
|
|
# Build graph
|
|
adj, order = build_correlation_graph(returns_dict, threshold=0.1)
|
|
print(f" Adjacency matrix:\n{adj}")
|
|
|
|
# Synthetic features and labels
|
|
features_dict = {c: np.random.randn(n, n_features).astype(np.float32) for c in coins}
|
|
labels_dict = {c: (np.random.rand(n) > 0.5).astype(np.float32) for c in coins}
|
|
|
|
model, history = train_gnn(
|
|
features_dict, labels_dict, adj, order,
|
|
n_features=n_features, epochs=5, batch_size=32
|
|
)
|
|
|
|
# Predict
|
|
current = {c: np.random.randn(n_features) for c in coins}
|
|
result = predict_gnn(model, current, adj, order)
|
|
for c, p in result.items():
|
|
print(f" {c}: {p:.4f}")
|
|
|
|
n_params = sum(p.numel() for p in model.parameters())
|
|
print(f" Parameters: {n_params:,}")
|
|
|
|
# Lead-lag graph
|
|
adj_ll, _ = build_lead_lag_graph(returns_dict, max_lag=3)
|
|
print(f" Lead-lag adjacency:\n{adj_ll}")
|
|
|
|
print("[GNN] Smoke test passed")
|