#!/usr/bin/env python3 """ Spatial-Temporal Transformer + GCN Model for Wuhan Respiratory Disease Risk Prediction. Architecture per PRD acceptance criteria: - Temporal Transformer: 3 layers, 4 heads - GCN: 2 layers [GCNConv(48, 128) → ReLU → Dropout(0.2) → GCNConv(128, 64)] - Input: [N, T, 48] node features, [N, N] adjacency - Output: [N, 3] risk values (1-day, 3-day, 7-day) """ import torch import torch.nn as nn import torch.nn.functional as F from torch_geometric.nn import GCNConv from torch_geometric.utils import add_self_loops class SpatialTemporalGCN(nn.Module): """ Spatial-Temporal Graph Convolutional Network with Transformer encoder. Args: node_features (int): Number of input node features (default: 48) temporal_heads (int): Number of attention heads in Transformer (default: 4) temporal_layers (int): Number of Transformer layers (default: 3) gcn_hidden (int): Hidden dimension for GCN layers (default: 128) gcn_output (int): Output dimension of GCN (default: 64) dropout (float): Dropout rate (default: 0.2) """ def __init__( self, node_features: int = 48, temporal_heads: int = 4, temporal_layers: int = 3, gcn_hidden: int = 128, gcn_output: int = 64, dropout: float = 0.2, ): super().__init__() # Temporal Transformer encoder encoder_layer = nn.TransformerEncoderLayer( d_model=node_features, nhead=temporal_heads, dim_feedforward=node_features * 4, dropout=dropout, activation='gelu', batch_first=True, norm_first=True, ) self.temporal_transformer = nn.TransformerEncoder( encoder_layer, num_layers=temporal_layers, ) # GCN layers self.conv1 = GCNConv(node_features, gcn_hidden) self.conv2 = GCNConv(gcn_hidden, gcn_output) self.dropout = nn.Dropout(dropout) self.relu = nn.ReLU() # Output head: 3 risk horizons (1-day, 3-day, 7-day) self.risk_head = nn.Linear(gcn_output, 3) def forward(self, x: torch.Tensor, edge_index: torch.Tensor) -> torch.Tensor: """ Forward pass. Args: x: Node features [N, T, 48] — N nodes, T time steps, 48 features edge_index: Graph connectivity [2, E] Returns: Risk predictions [N, 3] — 1-day, 3-day, 7-day risk """ N, T, F = x.shape # Temporal Transformer: process each node's time series # Input [N, T, 48] → Transformer → [N, T, 48] x_temporal = self.temporal_transformer(x) # Take the last time step as the spatial representation x_spatial = x_temporal[:, -1, :] # [N, 48] # Add self-loops for GCN edge_index, _ = add_self_loops(edge_index, num_nodes=N) # GCN layer 1: [N, 48] → [N, 128] x_gcn = self.conv1(x_spatial, edge_index) x_gcn = self.relu(x_gcn) x_gcn = self.dropout(x_gcn) # GCN layer 2: [N, 128] → [N, 64] x_gcn = self.conv2(x_gcn, edge_index) x_gcn = self.relu(x_gcn) x_gcn = self.dropout(x_gcn) # Risk prediction head: [N, 64] → [N, 3] risk = self.risk_head(x_gcn) # Clamp output to [0, 1] range (risk probability) risk = torch.sigmoid(risk) return risk def export_onnx(model, output_path: str, node_features: int = 48): """Export model to ONNX format for inference.""" model.eval() N = 512 # Dummy batch size for export # Dummy inputs matching expected shapes dummy_x = torch.randn(N, 14, node_features) # [N, T=14, 48] dummy_edge_index = torch.randint(0, N, (2, N * 4)) # Sparse edges torch.onnx.export( model, (dummy_x, dummy_edge_index), output_path, input_names=['node_features', 'edge_index'], output_names=['risk'], dynamic_axes={ 'node_features': {0: 'num_nodes'}, 'edge_index': {1: 'num_edges'}, 'risk': {0: 'num_nodes'}, }, opset_version=17, ) print(f"ONNX model exported to {output_path}") if __name__ == '__main__': # Quick forward pass test on dummy data model = SpatialTemporalGCN() # Dummy input: [512 nodes, 14 time steps, 48 features] N, T, F = 512, 14, 48 x = torch.randn(N, T, F) edge_index = torch.randint(0, N, (2, N * 4)) risk = model(x, edge_index) print(f"Input: {x.shape}") print(f"Edge index: {edge_index.shape}") print(f"Output risk: {risk.shape} — 1d:{risk[:,0].mean():.3f}, 3d:{risk[:,1].mean():.3f}, 7d:{risk[:,2].mean():.3f}") # ONNX export export_onnx(model, 'models/spatiotemporal_gcn/model_1_3_7.onnx')