//! Icosahedral mesh encoder for weather prediction. //! //! This module implements the mesh structure and encoding/decoding //! for converting between regular grid data and icosahedral mesh //! representation used in GraphCast-style weather models. use weathercast_shared::{AtmosphericState, MeshConfig}; // ============================================================================ // Icosahedral Mesh Structure // ============================================================================ /// Icosahedral mesh for global weather modeling. #[derive(Debug, Clone)] pub struct IcosahedralMesh { /// Mesh vertices as (x, y, z) on unit sphere. pub vertices: Vec<(f32, f32, f32)>, /// Mesh edges as vertex index pairs. pub edges: Vec<(usize, usize)>, /// Mesh faces as vertex index triplets. pub faces: Vec<(usize, usize, usize)>, /// Refinement level. pub refinement_level: usize, /// Average edge length in km. pub avg_edge_length: f32, } impl IcosahedralMesh { /// Create a new icosahedral mesh with given configuration. pub fn new(config: &MeshConfig) -> Self { let (vertices, edges, faces) = Self::generate_icosahedral_mesh(config.refinement_levels); Self { vertices, edges, faces, refinement_level: config.refinement_levels, avg_edge_length: config.avg_edge_length, } } /// Generate icosahedral mesh at given refinement level. fn generate_icosahedral_mesh( refinement: usize, ) -> ( Vec<(f32, f32, f32)>, Vec<(usize, usize)>, Vec<(usize, usize, usize)>, ) { // Start with base icosahedron (12 vertices, 20 faces, 30 edges) let phi = f32::midpoint(1.0, 5.0_f32.sqrt()); // Golden ratio let mut vertices = vec![ (-1.0, phi, 0.0), (1.0, phi, 0.0), (-1.0, -phi, 0.0), (1.0, -phi, 0.0), (0.0, -1.0, phi), (0.0, 1.0, phi), (0.0, -1.0, -phi), (0.0, 1.0, -phi), (phi, 0.0, -1.0), (phi, 0.0, 1.0), (-phi, 0.0, -1.0), (-phi, 0.0, 1.0), ]; // Normalize to unit sphere for v in &mut vertices { let len = (v.0 * v.0 + v.1 * v.1 + v.2 * v.2).sqrt(); v.0 /= len; v.1 /= len; v.2 /= len; } // Base faces (20 triangular faces of icosahedron) let mut faces = vec![ (0, 11, 5), (0, 5, 1), (0, 1, 7), (0, 7, 10), (0, 10, 11), (1, 5, 9), (5, 11, 4), (11, 10, 2), (10, 7, 6), (7, 1, 8), (3, 9, 4), (3, 4, 2), (3, 2, 6), (3, 6, 8), (3, 8, 9), (4, 9, 5), (2, 4, 11), (6, 2, 10), (8, 6, 7), (9, 8, 1), ]; // Refine mesh for _ in 0..refinement { let (new_vertices, new_faces) = Self::subdivide(&vertices, &faces); vertices = new_vertices; faces = new_faces; } // Generate edges from faces let mut edges = Vec::new(); let mut edge_set = std::collections::HashSet::new(); for &(a, b, c) in &faces { for (v1, v2) in [(a, b), (b, c), (c, a)] { let edge = if v1 < v2 { (v1, v2) } else { (v2, v1) }; if edge_set.insert(edge) { edges.push(edge); } } } (vertices, edges, faces) } /// Subdivide mesh faces (each triangle becomes 4 triangles). fn subdivide( vertices: &[(f32, f32, f32)], faces: &[(usize, usize, usize)], ) -> (Vec<(f32, f32, f32)>, Vec<(usize, usize, usize)>) { let mut new_vertices = vertices.to_vec(); let mut new_faces = Vec::new(); let mut midpoint_cache: std::collections::HashMap<(usize, usize), usize> = std::collections::HashMap::new(); let get_midpoint = |v1: usize, v2: usize, vertices: &[(f32, f32, f32)], new_vertices: &mut Vec<(f32, f32, f32)>, cache: &mut std::collections::HashMap<(usize, usize), usize>| -> usize { let key = if v1 < v2 { (v1, v2) } else { (v2, v1) }; if let Some(&idx) = cache.get(&key) { return idx; } let p1 = vertices[v1]; let p2 = vertices[v2]; // Calculate midpoint and project to sphere let mut mid = ( f32::midpoint(p1.0, p2.0), f32::midpoint(p1.1, p2.1), f32::midpoint(p1.2, p2.2), ); let len = (mid.0 * mid.0 + mid.1 * mid.1 + mid.2 * mid.2).sqrt(); mid.0 /= len; mid.1 /= len; mid.2 /= len; let idx = new_vertices.len(); new_vertices.push(mid); cache.insert(key, idx); idx }; for &(a, b, c) in faces { let ab = get_midpoint(a, b, vertices, &mut new_vertices, &mut midpoint_cache); let bc = get_midpoint(b, c, vertices, &mut new_vertices, &mut midpoint_cache); let ca = get_midpoint(c, a, vertices, &mut new_vertices, &mut midpoint_cache); new_faces.push((a, ab, ca)); new_faces.push((b, bc, ab)); new_faces.push((c, ca, bc)); new_faces.push((ab, bc, ca)); } (new_vertices, new_faces) } /// Get vertex latitude/longitude. pub fn vertex_latlon(&self, idx: usize) -> Option<(f32, f32)> { self.vertices.get(idx).map(|&(x, y, z)| { let lat = y.asin().to_degrees(); let lon = x.atan2(z).to_degrees(); (lat, lon) }) } /// Find nearest vertex to given lat/lon. pub fn nearest_vertex(&self, lat: f32, lon: f32) -> usize { let lat_rad = lat.to_radians(); let lon_rad = lon.to_radians(); // Convert to Cartesian let x = lat_rad.cos() * lon_rad.sin(); let y = lat_rad.sin(); let z = lat_rad.cos() * lon_rad.cos(); // Find nearest vertex self.vertices .iter() .enumerate() .min_by(|(_, v1), (_, v2)| { let d1 = (v1.0 - x).powi(2) + (v1.1 - y).powi(2) + (v1.2 - z).powi(2); let d2 = (v2.0 - x).powi(2) + (v2.1 - y).powi(2) + (v2.2 - z).powi(2); d1.partial_cmp(&d2).unwrap_or(std::cmp::Ordering::Equal) }) .map_or(0, |(i, _)| i) } /// Get neighbors of a vertex. pub fn vertex_neighbors(&self, idx: usize) -> Vec { self.edges .iter() .filter_map(|&(a, b)| { if a == idx { Some(b) } else if b == idx { Some(a) } else { None } }) .collect() } } // ============================================================================ // Mesh Encoder // ============================================================================ /// Encoder for converting between grid and mesh representations. #[derive(Debug)] pub struct MeshEncoder { /// Number of mesh nodes. #[allow(dead_code)] num_nodes: usize, /// Hidden dimension. hidden_dim: usize, /// Encoder weights. encoder_weights: Vec, /// Decoder weights. decoder_weights: Vec, /// RNG state. rng_state: u64, } impl MeshEncoder { /// Create a new mesh encoder. pub fn new(num_nodes: usize, hidden_dim: usize) -> Self { let mut encoder = Self { num_nodes, hidden_dim, encoder_weights: Vec::new(), decoder_weights: Vec::new(), rng_state: 42, }; encoder.initialize_weights(); encoder } /// Initialize encoder/decoder weights. fn initialize_weights(&mut self) { // Encoder: input_dim -> hidden_dim let input_dim = 8 * 5 + 5; // 8 pressure levels * 5 vars + 5 surface vars self.encoder_weights = (0..input_dim * self.hidden_dim) .map(|_| self.random() as f32 * 0.1 - 0.05) .collect(); // Decoder: hidden_dim -> output_dim self.decoder_weights = (0..self.hidden_dim * input_dim) .map(|_| self.random() as f32 * 0.1 - 0.05) .collect(); } /// Encode atmospheric state to mesh representation. pub fn encode_state(&mut self, state: &AtmosphericState, mesh: &IcosahedralMesh) -> Vec { let num_nodes = mesh.vertices.len(); let mut mesh_state = vec![0.0f32; num_nodes * self.hidden_dim]; // Flatten state to feature vector let features = self.state_to_features(state); // Apply encoder at each mesh node for node in 0..num_nodes { let (lat, lon) = mesh.vertex_latlon(node).unwrap_or((0.0, 0.0)); // Modulate features by position let lat_factor = lat.to_radians().cos(); let lon_factor = (lon.to_radians() * 2.0).sin() * 0.1; for h in 0..self.hidden_dim { let mut val = 0.0f32; // Linear projection with positional modulation for (f_idx, &f) in features.iter().enumerate() { let w_idx = f_idx * self.hidden_dim + h; let weight = self.encoder_weights.get(w_idx).copied().unwrap_or(0.01); val += f * weight; } // Add positional encoding val += (lat_factor * (h as f32 * 0.1)).sin() * 0.5; val += lon_factor * (h as f32 * 0.1).cos() * 0.3; // Apply non-linearity mesh_state[node * self.hidden_dim + h] = self.gelu(val); } } mesh_state } /// Decode mesh representation back to atmospheric state. pub fn decode_state( &mut self, mesh_state: &[f32], mesh: &IcosahedralMesh, timestamp: f32, ) -> AtmosphericState { let num_nodes = mesh.vertices.len(); // Average over all nodes (simplified global mean) let mut avg_hidden = vec![0.0f32; self.hidden_dim]; for node in 0..num_nodes { for h in 0..self.hidden_dim { let idx = node * self.hidden_dim + h; avg_hidden[h] += mesh_state.get(idx).copied().unwrap_or(0.0) / num_nodes as f32; } } // Decode to output features let output_dim = 8 * 5 + 5; let mut output = vec![0.0f32; output_dim]; for o in 0..output_dim { for h in 0..self.hidden_dim { let w_idx = h * output_dim + o; let weight = self.decoder_weights.get(w_idx).copied().unwrap_or(0.01); output[o] += avg_hidden[h] * weight; } } // Reconstruct atmospheric state with physical constraints self.features_to_state(&output, timestamp) } /// Convert atmospheric state to feature vector. fn state_to_features(&self, state: &AtmosphericState) -> Vec { let mut features = Vec::new(); // Normalize and add pressure level variables for i in 0..8 { features.push(state.temperature.get(i).copied().unwrap_or(250.0) / 300.0); features.push(state.geopotential.get(i).copied().unwrap_or(5000.0) / 10000.0); features.push(state.humidity.get(i).copied().unwrap_or(0.005) * 100.0); features.push(state.wind_u.get(i).copied().unwrap_or(0.0) / 50.0); features.push(state.wind_v.get(i).copied().unwrap_or(0.0) / 50.0); } // Surface variables features.push(state.surface_pressure / 101325.0); features.push(state.temperature_2m / 300.0); features.push(state.wind_u_10m / 20.0); features.push(state.wind_v_10m / 20.0); features.push(state.precipitation / 10.0); features } /// Convert feature vector back to atmospheric state. fn features_to_state(&mut self, features: &[f32], timestamp: f32) -> AtmosphericState { let pressure_levels = vec![1000.0, 925.0, 850.0, 700.0, 500.0, 300.0, 200.0, 50.0]; let mut temperature = Vec::new(); let mut geopotential = Vec::new(); let mut humidity = Vec::new(); let mut wind_u = Vec::new(); let mut wind_v = Vec::new(); for i in 0..8 { let base = i * 5; // Clamp temperature to physically reasonable range (150K to 350K) let temp = (features.get(base).copied().unwrap_or(0.83) * 300.0).clamp(150.0, 350.0); temperature.push(temp); geopotential.push(features.get(base + 1).copied().unwrap_or(0.5) * 10000.0); humidity.push((features.get(base + 2).copied().unwrap_or(0.5) / 100.0).max(0.0)); wind_u.push(features.get(base + 3).copied().unwrap_or(0.0) * 50.0); wind_v.push(features.get(base + 4).copied().unwrap_or(0.0) * 50.0); } // Add some stochastic variability (small enough to stay in range) for t in &mut temperature { *t += (self.random() as f32 - 0.5) * 2.0; *t = t.clamp(150.0, 350.0); // Ensure still in range after noise } AtmosphericState { temperature, geopotential, humidity, wind_u, wind_v, pressure_levels, surface_pressure: features.get(40).copied().unwrap_or(1.0) * 101325.0, temperature_2m: (features.get(41).copied().unwrap_or(0.96) * 300.0).clamp(150.0, 350.0), wind_u_10m: features.get(42).copied().unwrap_or(0.0) * 20.0, wind_v_10m: features.get(43).copied().unwrap_or(0.0) * 20.0, precipitation: (features.get(44).copied().unwrap_or(0.0) * 10.0).max(0.0), timestamp, } } /// GELU activation function. 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::*; #[test] fn test_mesh_creation() { let config = MeshConfig::from_refinement(3); let mesh = IcosahedralMesh::new(&config); assert!(!mesh.vertices.is_empty()); assert!(!mesh.edges.is_empty()); assert!(!mesh.faces.is_empty()); } #[test] fn test_mesh_refinement_levels() { // Test different refinement levels for level in 0..4 { let config = MeshConfig::from_refinement(level); let mesh = IcosahedralMesh::new(&config); // Icosahedral mesh: V = 10 * 4^n + 2 let expected_vertices = 10 * 4_usize.pow(level as u32) + 2; assert_eq!(mesh.vertices.len(), expected_vertices); } } #[test] fn test_vertices_on_unit_sphere() { let config = MeshConfig::from_refinement(2); let mesh = IcosahedralMesh::new(&config); for &(x, y, z) in &mesh.vertices { let len = (x * x + y * y + z * z).sqrt(); assert!((len - 1.0).abs() < 0.001, "Vertex not on unit sphere"); } } #[test] fn test_vertex_latlon() { let config = MeshConfig::from_refinement(2); let mesh = IcosahedralMesh::new(&config); for idx in 0..mesh.vertices.len() { let latlon = mesh.vertex_latlon(idx); assert!(latlon.is_some()); let (lat, lon) = latlon.unwrap(); assert!(lat >= -90.0 && lat <= 90.0); assert!(lon >= -180.0 && lon <= 180.0); } } #[test] fn test_nearest_vertex() { let config = MeshConfig::from_refinement(3); let mesh = IcosahedralMesh::new(&config); // Find vertex near New York City let nyc_vertex = mesh.nearest_vertex(40.7, -74.0); let (lat, _lon) = mesh.vertex_latlon(nyc_vertex).unwrap(); // Should be reasonably close assert!((lat - 40.7).abs() < 10.0); } #[test] fn test_vertex_neighbors() { let config = MeshConfig::from_refinement(2); let mesh = IcosahedralMesh::new(&config); let neighbors = mesh.vertex_neighbors(0); // Each vertex should have multiple neighbors (typically 5-6 for icosahedral mesh) assert!(!neighbors.is_empty()); assert!(neighbors.len() >= 3); } #[test] fn test_mesh_encoder_creation() { let encoder = MeshEncoder::new(100, 64); assert_eq!(encoder.num_nodes, 100); assert_eq!(encoder.hidden_dim, 64); } #[test] fn test_encode_decode_roundtrip() { let config = MeshConfig::from_refinement(2); let mesh = IcosahedralMesh::new(&config); let mut encoder = MeshEncoder::new(mesh.vertices.len(), 64); let original_state = weathercast_shared::sample_atmospheric_state(); let encoded = encoder.encode_state(&original_state, &mesh); let decoded = encoder.decode_state(&encoded, &mesh, 0.0); // Check that decoded state has same structure assert_eq!(decoded.temperature.len(), original_state.temperature.len()); assert_eq!( decoded.pressure_levels.len(), original_state.pressure_levels.len() ); // Values should be in reasonable ranges for t in &decoded.temperature { assert!(*t > 100.0 && *t < 400.0, "Temperature out of range: {}", t); } } #[test] fn test_encoder_mesh_state_size() { let config = MeshConfig::from_refinement(2); let mesh = IcosahedralMesh::new(&config); let hidden_dim = 64; let mut encoder = MeshEncoder::new(mesh.vertices.len(), hidden_dim); let state = weathercast_shared::sample_atmospheric_state(); let encoded = encoder.encode_state(&state, &mesh); assert_eq!(encoded.len(), mesh.vertices.len() * hidden_dim); } #[test] fn test_gelu_activation() { let encoder = MeshEncoder::new(10, 32); // GELU(0) should be approximately 0 assert!(encoder.gelu(0.0).abs() < 0.01); // GELU should be monotonically increasing for positive x assert!(encoder.gelu(1.0) > encoder.gelu(0.5)); assert!(encoder.gelu(0.5) > encoder.gelu(0.0)); } }