//! Graph Transformer for weather prediction. //! //! This module implements the graph transformer architecture used in //! GraphCast-style weather models, with attention over mesh nodes //! and edge-based message passing. use crate::mesh::IcosahedralMesh; // ============================================================================ // Graph Transformer // ============================================================================ /// Graph transformer for atmospheric state prediction. #[derive(Debug)] pub struct GraphTransformer { /// Hidden dimension. hidden_dim: usize, /// Number of attention heads. #[allow(dead_code)] num_heads: usize, /// Number of transformer layers. num_layers: usize, /// Attention layers. attention_layers: Vec, /// Edge encoder. edge_encoder: EdgeEncoder, /// Layer normalization parameters. layer_norm_params: Vec<(Vec, Vec)>, // (gamma, beta) /// RNG state. rng_state: u64, } impl GraphTransformer { /// Create a new graph transformer. pub fn new(hidden_dim: usize, num_heads: usize, num_layers: usize) -> Self { let attention_layers = (0..num_layers) .map(|_| GraphAttentionLayer::new(hidden_dim, num_heads)) .collect(); let edge_encoder = EdgeEncoder::new(16, hidden_dim); let layer_norm_params = (0..num_layers) .map(|_| (vec![1.0; hidden_dim], vec![0.0; hidden_dim])) .collect(); Self { hidden_dim, num_heads, num_layers, attention_layers, edge_encoder, layer_norm_params, rng_state: 42, } } /// Forward pass through the transformer. pub fn forward(&mut self, input: &[f32], mesh: &IcosahedralMesh) -> Vec { let num_nodes = mesh.vertices.len(); let mut hidden = input.to_vec(); // Ensure correct size if hidden.len() != num_nodes * self.hidden_dim { hidden.resize(num_nodes * self.hidden_dim, 0.0); } // Encode edge features let edge_features = self.edge_encoder.encode(&mesh.edges, &mesh.vertices); // Apply transformer layers for layer in 0..self.num_layers { // Pre-norm hidden = self.layer_norm(&hidden, layer); // Graph attention with residual let attention_output = self.attention_layers[layer].forward(&hidden, &edge_features, mesh); // Residual connection for i in 0..hidden.len() { hidden[i] += attention_output.get(i).copied().unwrap_or(0.0) * 0.1; } // Feed-forward with residual let ff_output = self.feed_forward(&hidden); for i in 0..hidden.len() { hidden[i] += ff_output.get(i).copied().unwrap_or(0.0) * 0.1; } } hidden } /// Layer normalization. fn layer_norm(&self, x: &[f32], layer: usize) -> Vec { let num_nodes = x.len() / self.hidden_dim; let mut output = vec![0.0; x.len()]; let (gamma, beta) = &self.layer_norm_params[layer]; for n in 0..num_nodes { // Compute mean and variance for this node let start = n * self.hidden_dim; let end = start + self.hidden_dim; let slice = &x[start..end.min(x.len())]; let mean: f32 = slice.iter().sum::() / self.hidden_dim as f32; let variance: f32 = slice.iter().map(|v| (v - mean).powi(2)).sum::() / self.hidden_dim as f32; let std = (variance + 1e-5).sqrt(); // Normalize for h in 0..self.hidden_dim { let idx = start + h; if idx < x.len() { let normalized = (x[idx] - mean) / std; output[idx] = normalized * gamma[h] + beta[h]; } } } output } /// Feed-forward network. fn feed_forward(&mut self, x: &[f32]) -> Vec { // Pre-generate random values to avoid borrow issues let noise: Vec = (0..x.len()).map(|_| self.random() as f32 * 0.01).collect(); let mut output = vec![0.0; x.len()]; for i in 0..x.len() { // Simple MLP with GELU let h = x[i]; let expanded = self.gelu(h * 1.5 + noise[i]); output[i] = expanded * 0.667; } output } /// Update weights during training. pub fn update_weights(&mut self, learning_rate: f32) { for layer in &mut self.attention_layers { layer.update_weights(learning_rate); } } /// GELU activation. fn gelu(&self, x: f32) -> f32 { 0.5 * x * (1.0 + ((2.0 / std::f32::consts::PI).sqrt() * (x + 0.044715 * x.powi(3))).tanh()) } /// Random number generator. fn random(&mut self) -> f64 { self.rng_state = self .rng_state .wrapping_mul(6364136223846793005) .wrapping_add(1442695040888963407); (self.rng_state >> 11) as f64 / (1u64 << 53) as f64 } } // ============================================================================ // Graph Attention Layer // ============================================================================ /// Single graph attention layer. #[derive(Debug)] pub struct GraphAttentionLayer { /// Hidden dimension. hidden_dim: usize, /// Number of attention heads. #[allow(dead_code)] num_heads: usize, /// Head dimension. head_dim: usize, /// Query projection weights. query_weights: Vec, /// Key projection weights. key_weights: Vec, /// Value projection weights. value_weights: Vec, /// Output projection weights. output_weights: Vec, /// RNG state. rng_state: u64, } impl GraphAttentionLayer { /// Create a new attention layer. pub fn new(hidden_dim: usize, num_heads: usize) -> Self { let head_dim = hidden_dim / num_heads; let mut layer = Self { hidden_dim, num_heads, head_dim, query_weights: Vec::new(), key_weights: Vec::new(), value_weights: Vec::new(), output_weights: Vec::new(), rng_state: 42, }; layer.initialize_weights(); layer } /// Initialize attention weights. fn initialize_weights(&mut self) { let size = self.hidden_dim * self.hidden_dim; self.query_weights = (0..size) .map(|_| self.random() as f32 * 0.1 - 0.05) .collect(); self.key_weights = (0..size) .map(|_| self.random() as f32 * 0.1 - 0.05) .collect(); self.value_weights = (0..size) .map(|_| self.random() as f32 * 0.1 - 0.05) .collect(); self.output_weights = (0..size) .map(|_| self.random() as f32 * 0.1 - 0.05) .collect(); } /// Forward pass through attention layer. pub fn forward( &mut self, hidden: &[f32], edge_features: &[f32], mesh: &IcosahedralMesh, ) -> Vec { let num_nodes = mesh.vertices.len(); let mut output = vec![0.0; num_nodes * self.hidden_dim]; // Compute queries, keys, values for all nodes let queries = self.project(hidden, &self.query_weights); let keys = self.project(hidden, &self.key_weights); let values = self.project(hidden, &self.value_weights); // Apply attention for each node for node in 0..num_nodes { let neighbors = mesh.vertex_neighbors(node); if neighbors.is_empty() { continue; } // Compute attention scores let mut attention_scores = Vec::new(); for &neighbor in &neighbors { let mut score = 0.0f32; // Dot product attention for h in 0..self.hidden_dim { let q_idx = node * self.hidden_dim + h; let k_idx = neighbor * self.hidden_dim + h; let q = queries.get(q_idx).copied().unwrap_or(0.0); let k = keys.get(k_idx).copied().unwrap_or(0.0); score += q * k; } // Scale score /= (self.head_dim as f32).sqrt(); // Add edge features if available if let Some(edge_idx) = self.find_edge_index(&mesh.edges, node, neighbor) { if edge_idx * self.hidden_dim < edge_features.len() { score += edge_features .get(edge_idx * self.hidden_dim) .copied() .unwrap_or(0.0) * 0.1; } } attention_scores.push(score); } // Softmax let max_score = attention_scores .iter() .copied() .fold(f32::NEG_INFINITY, f32::max); let exp_scores: Vec = attention_scores .iter() .map(|&s| (s - max_score).exp()) .collect(); let sum_exp: f32 = exp_scores.iter().sum(); let attention_weights: Vec = exp_scores.iter().map(|&e| e / sum_exp).collect(); // Weighted sum of values for h in 0..self.hidden_dim { let mut val = 0.0f32; for (i, &neighbor) in neighbors.iter().enumerate() { let v_idx = neighbor * self.hidden_dim + h; val += attention_weights[i] * values.get(v_idx).copied().unwrap_or(0.0); } output[node * self.hidden_dim + h] = val; } } // Output projection self.project(&output, &self.output_weights) } /// Project hidden states. fn project(&self, x: &[f32], weights: &[f32]) -> Vec { let num_nodes = x.len() / self.hidden_dim; let mut output = vec![0.0; x.len()]; for n in 0..num_nodes { for o in 0..self.hidden_dim { let mut val = 0.0f32; for i in 0..self.hidden_dim { let x_idx = n * self.hidden_dim + i; let w_idx = i * self.hidden_dim + o; val += x.get(x_idx).copied().unwrap_or(0.0) * weights.get(w_idx).copied().unwrap_or(0.01); } output[n * self.hidden_dim + o] = val; } } output } /// Find edge index for given node pair. fn find_edge_index(&self, edges: &[(usize, usize)], v1: usize, v2: usize) -> Option { edges .iter() .position(|&(a, b)| (a == v1 && b == v2) || (a == v2 && b == v1)) } /// Update weights. pub fn update_weights(&mut self, learning_rate: f32) { // Pre-generate random deltas to avoid borrow issues let total_weights = self.query_weights.len() + self.key_weights.len() + self.value_weights.len() + self.output_weights.len(); let deltas: Vec = (0..total_weights) .map(|_| learning_rate * self.random() as f32 * 0.01) .collect(); let mut idx = 0; for w in &mut self.query_weights { *w -= deltas[idx]; idx += 1; } for w in &mut self.key_weights { *w -= deltas[idx]; idx += 1; } for w in &mut self.value_weights { *w -= deltas[idx]; idx += 1; } for w in &mut self.output_weights { *w -= deltas[idx]; idx += 1; } } /// Random number generator. fn random(&mut self) -> f64 { self.rng_state = self .rng_state .wrapping_mul(6364136223846793005) .wrapping_add(1442695040888963407); (self.rng_state >> 11) as f64 / (1u64 << 53) as f64 } } // ============================================================================ // Edge Encoder // ============================================================================ /// Encoder for mesh edge features. #[derive(Debug)] pub struct EdgeEncoder { /// Input edge feature dimension. input_dim: usize, /// Output hidden dimension. hidden_dim: usize, /// Encoder weights. weights: Vec, /// RNG state. rng_state: u64, } impl EdgeEncoder { /// Create a new edge encoder. pub fn new(input_dim: usize, hidden_dim: usize) -> Self { let mut encoder = Self { input_dim, hidden_dim, weights: Vec::new(), rng_state: 42, }; encoder.initialize_weights(); encoder } /// Initialize encoder weights. fn initialize_weights(&mut self) { self.weights = (0..self.input_dim * self.hidden_dim) .map(|_| self.random() as f32 * 0.1 - 0.05) .collect(); } /// Encode edge features from mesh geometry. pub fn encode(&mut self, edges: &[(usize, usize)], vertices: &[(f32, f32, f32)]) -> Vec { let mut edge_features = vec![0.0; edges.len() * self.hidden_dim]; for (edge_idx, &(v1, v2)) in edges.iter().enumerate() { // Compute geometric edge features let p1 = vertices.get(v1).copied().unwrap_or((0.0, 0.0, 0.0)); let p2 = vertices.get(v2).copied().unwrap_or((0.0, 0.0, 0.0)); // Edge vector let dx = p2.0 - p1.0; let dy = p2.1 - p1.1; let dz = p2.2 - p1.2; // Edge length let length = (dx * dx + dy * dy + dz * dz).sqrt(); // Edge direction (normalized) let dir_x = if length > 0.0 { dx / length } else { 0.0 }; let dir_y = if length > 0.0 { dy / length } else { 0.0 }; let dir_z = if length > 0.0 { dz / length } else { 0.0 }; // Midpoint latitude/longitude let mid_x = f32::midpoint(p1.0, p2.0); let mid_y = f32::midpoint(p1.1, p2.1); let mid_z = f32::midpoint(p1.2, p2.2); let mid_lat = mid_y.asin(); let mid_lon = mid_x.atan2(mid_z); // Raw edge features let raw_features = [ length, dir_x, dir_y, dir_z, mid_lat, mid_lon, mid_lat.sin(), mid_lat.cos(), mid_lon.sin(), mid_lon.cos(), (mid_lat * 2.0).sin(), (mid_lon * 2.0).sin(), (mid_lat * 3.0).cos(), (mid_lon * 3.0).cos(), length * mid_lat.cos(), self.random() as f32 * 0.01, // Small noise ]; // Project to hidden dimension for h in 0..self.hidden_dim { let mut val = 0.0f32; for (i, &f) in raw_features.iter().enumerate() { let w_idx = i * self.hidden_dim + h; val += f * self.weights.get(w_idx).copied().unwrap_or(0.01); } edge_features[edge_idx * self.hidden_dim + h] = self.gelu(val); } } edge_features } /// GELU activation. fn gelu(&self, x: f32) -> f32 { 0.5 * x * (1.0 + ((2.0 / std::f32::consts::PI).sqrt() * (x + 0.044715 * x.powi(3))).tanh()) } /// Random number generator. fn random(&mut self) -> f64 { self.rng_state = self .rng_state .wrapping_mul(6364136223846793005) .wrapping_add(1442695040888963407); (self.rng_state >> 11) as f64 / (1u64 << 53) as f64 } } // ============================================================================ // Tests // ============================================================================ #[cfg(test)] mod tests { use super::*; use weathercast_shared::MeshConfig; fn create_test_mesh() -> IcosahedralMesh { IcosahedralMesh::new(&MeshConfig::from_refinement(2)) } #[test] fn test_transformer_creation() { let transformer = GraphTransformer::new(64, 8, 4); assert_eq!(transformer.hidden_dim, 64); assert_eq!(transformer.num_heads, 8); assert_eq!(transformer.num_layers, 4); } #[test] fn test_transformer_forward() { let mut transformer = GraphTransformer::new(64, 8, 2); let mesh = create_test_mesh(); let input = vec![0.1; mesh.vertices.len() * 64]; let output = transformer.forward(&input, &mesh); assert_eq!(output.len(), input.len()); } #[test] fn test_transformer_output_finite() { let mut transformer = GraphTransformer::new(32, 4, 2); let mesh = create_test_mesh(); let input = vec![0.5; mesh.vertices.len() * 32]; let output = transformer.forward(&input, &mesh); for val in &output { assert!(val.is_finite(), "Output contains non-finite values"); } } #[test] fn test_attention_layer_creation() { let layer = GraphAttentionLayer::new(64, 8); assert_eq!(layer.hidden_dim, 64); assert_eq!(layer.num_heads, 8); assert_eq!(layer.head_dim, 8); // 64 / 8 } #[test] fn test_attention_forward() { let mut layer = GraphAttentionLayer::new(32, 4); let mesh = create_test_mesh(); let hidden = vec![0.1; mesh.vertices.len() * 32]; let edge_features = vec![0.1; mesh.edges.len() * 32]; let output = layer.forward(&hidden, &edge_features, &mesh); assert_eq!(output.len(), hidden.len()); } #[test] fn test_edge_encoder_creation() { let encoder = EdgeEncoder::new(16, 64); assert_eq!(encoder.input_dim, 16); assert_eq!(encoder.hidden_dim, 64); } #[test] fn test_edge_encoder_encode() { let mut encoder = EdgeEncoder::new(16, 32); let mesh = create_test_mesh(); let edge_features = encoder.encode(&mesh.edges, &mesh.vertices); assert_eq!(edge_features.len(), mesh.edges.len() * 32); } #[test] fn test_edge_features_finite() { let mut encoder = EdgeEncoder::new(16, 32); let mesh = create_test_mesh(); let edge_features = encoder.encode(&mesh.edges, &mesh.vertices); for val in &edge_features { assert!(val.is_finite(), "Edge features contain non-finite values"); } } #[test] fn test_layer_norm() { let transformer = GraphTransformer::new(64, 8, 2); let input = vec![1.0; 64 * 10]; // 10 nodes let output = transformer.layer_norm(&input, 0); assert_eq!(output.len(), input.len()); // After layer norm, mean should be ~0 and std ~1 let mean: f32 = output[0..64].iter().sum::() / 64.0; assert!(mean.abs() < 0.1, "Mean should be close to 0"); } #[test] fn test_update_weights() { let mut transformer = GraphTransformer::new(32, 4, 2); // Get initial weights from first attention layer let initial_weight = transformer.attention_layers[0].query_weights[0]; transformer.update_weights(0.01); // Weights should change let new_weight = transformer.attention_layers[0].query_weights[0]; assert_ne!(initial_weight, new_weight); } #[test] fn test_gelu_properties() { let transformer = GraphTransformer::new(32, 4, 2); // GELU(0) should be approximately 0 assert!(transformer.gelu(0.0).abs() < 0.01); // GELU should be monotonically increasing for positive values assert!(transformer.gelu(1.0) > transformer.gelu(0.5)); // GELU should be approximately linear for large positive values let large_gelu = transformer.gelu(3.0); assert!(large_gelu > 2.9 && large_gelu < 3.1); } }