//! WeatherCast: Regional Weather Prediction with GraphCast-style Architecture //! //! This demo showcases a GraphCast-inspired weather prediction system using //! graph neural networks on an icosahedral mesh. It demonstrates: //! - Graph transformer architecture for atmospheric state prediction //! - Icosahedral mesh encoding/decoding //! - Multi-scale mesh processing //! - Ensemble prediction and uncertainty quantification //! - Autoregressive multi-step forecasting pub mod mesh; pub mod sample_data; pub mod transformer; use mesh::{IcosahedralMesh, MeshEncoder}; use transformer::GraphTransformer; use weathercast_shared::{ AtmosphericState, GridPoint, MeshConfig, PredictionConfig, TrainingConfig, TrainingProgress, UncertaintyEstimate, VerificationMetrics, WeatherForecast, }; // ============================================================================ // Main WeatherCast System // ============================================================================ /// Main WeatherCast system for weather prediction. #[derive(Debug)] pub struct WeatherCast { /// Icosahedral mesh structure. mesh: IcosahedralMesh, /// Mesh encoder for grid-to-mesh conversion. mesh_encoder: MeshEncoder, /// Graph transformer for prediction. transformer: GraphTransformer, /// Prediction configuration. config: PredictionConfig, /// Whether model is trained. trained: bool, /// RNG state. rng_state: u64, } impl WeatherCast { /// Create a new WeatherCast system. pub fn new(mesh_config: MeshConfig, prediction_config: PredictionConfig) -> Self { let mesh = IcosahedralMesh::new(&mesh_config); let mesh_encoder = MeshEncoder::new(mesh_config.num_nodes, 128); let transformer = GraphTransformer::new(128, 8, 6); Self { mesh, mesh_encoder, transformer, config: prediction_config, trained: false, rng_state: 42, } } /// Create with default configuration. pub fn default_config() -> Self { Self::new(MeshConfig::default(), PredictionConfig::default()) } /// Generate weather forecast from initial state. pub fn forecast(&mut self, initial_state: &AtmosphericState) -> WeatherForecast { let start = std::time::Instant::now(); // Encode initial state to mesh representation let mesh_state = self.mesh_encoder.encode_state(initial_state, &self.mesh); // Run autoregressive prediction let mut states = Vec::new(); let mut lead_times = Vec::new(); let mut current_state = mesh_state; for step in 0..self.config.num_steps { let lead_time = (step as f32 + 1.0) * self.config.time_step; // Apply graph transformer current_state = self.transformer.forward(¤t_state, &self.mesh); // Decode mesh state back to atmospheric state let atm_state = self .mesh_encoder .decode_state(¤t_state, &self.mesh, lead_time); states.push(atm_state); lead_times.push(lead_time); } let computation_time_ms = start.elapsed().as_secs_f64() * 1000.0; // Generate grid point locations from mesh let locations: Vec = self .mesh .vertices .iter() .take(10) // Sample locations .map(|v| { let lat = v.1.asin().to_degrees(); let lon = v.0.atan2(v.2).to_degrees(); GridPoint::new(lat, lon) }) .collect(); WeatherForecast { init_time: format!( "2024-01-01T00:00:00Z (computed in {:.1}ms)", computation_time_ms ), hourly_states: states, lead_times, locations, ensemble_member: None, model_version: "weathercast-v1.0".to_string(), } } /// Generate ensemble forecast. pub fn ensemble_forecast( &mut self, initial_state: &AtmosphericState, num_members: usize, ) -> Vec { let mut forecasts = Vec::new(); for member in 0..num_members { // Perturb initial conditions let mut perturbed = initial_state.clone(); for t in &mut perturbed.temperature { *t += self.random() as f32 * 0.5 - 0.25; } for u in &mut perturbed.wind_u { *u += self.random() as f32 * 0.5 - 0.25; } let mut forecast = self.forecast(&perturbed); forecast.ensemble_member = Some(member); forecasts.push(forecast); } forecasts } /// Verify forecast against observations. pub fn verify( &self, forecast: &WeatherForecast, observations: &[AtmosphericState], ) -> Vec { let mut metrics = Vec::new(); // Verify at each lead time where we have observations for (idx, state) in forecast.hourly_states.iter().enumerate() { if idx >= observations.len() { break; } let obs = &observations[idx]; let lead_time = forecast.lead_times.get(idx).copied().unwrap_or(0.0); // Verify 500 hPa geopotential (key metric) let level_500_idx = state .pressure_levels .iter() .position(|&p| (p - 500.0).abs() < 1.0); if let Some(idx500) = level_500_idx { let forecast_z = &[state.geopotential.get(idx500).copied().unwrap_or(0.0)]; let obs_z = &[obs.geopotential.get(idx500).copied().unwrap_or(0.0)]; let mut m = VerificationMetrics::calculate(forecast_z, obs_z, "geopotential", lead_time); m.pressure_level = Some(500.0); metrics.push(m); } // Verify temperature let temp_metrics = VerificationMetrics::calculate( &state.temperature, &obs.temperature, "temperature", lead_time, ); metrics.push(temp_metrics); // Verify wind let wind_metrics = VerificationMetrics::calculate(&state.wind_u, &obs.wind_u, "wind_u", lead_time); metrics.push(wind_metrics); } metrics } /// Calculate uncertainty from ensemble. pub fn calculate_uncertainty( &self, ensemble: &[WeatherForecast], lead_time_idx: usize, variable: &str, ) -> Vec { if ensemble.is_empty() { return Vec::new(); } let num_levels = ensemble[0] .hourly_states .first() .map_or(0, |s| s.temperature.len()); (0..num_levels) .map(|level| { let values: Vec = ensemble .iter() .filter_map(|f| f.hourly_states.get(lead_time_idx)) .filter_map(|s| match variable { "temperature" => s.temperature.get(level).copied(), "geopotential" => s.geopotential.get(level).copied(), "humidity" => s.humidity.get(level).copied(), "wind_u" => s.wind_u.get(level).copied(), "wind_v" => s.wind_v.get(level).copied(), _ => None, }) .collect(); UncertaintyEstimate::from_ensemble(&values) }) .collect() } /// Train the model. pub fn train( &mut self, config: &TrainingConfig, progress_callback: Option>, ) { for epoch in 0..config.epochs { // Simulated training let train_loss = 0.5 * (-(epoch as f32) / 30.0).exp() + 0.05; let val_loss = train_loss * 1.1; // Update transformer weights (simulated) self.transformer.update_weights(config.learning_rate); // Key metrics let rmse_500hpa_z = 20.0 * (-(epoch as f32) / 50.0).exp() + 15.0; let acc_500hpa_z = 0.6 + 0.35 * (1.0 - (-(epoch as f32) / 40.0).exp()); if let Some(ref callback) = progress_callback { callback(TrainingProgress { epoch: epoch + 1, total_epochs: config.epochs, train_loss, val_loss, learning_rate: config.learning_rate, rmse_500hpa_z, acc_500hpa_z, }); } } self.trained = true; } /// Check if model is trained. pub fn is_trained(&self) -> bool { self.trained } /// Get mesh statistics. pub fn mesh_stats(&self) -> (usize, usize, usize) { ( self.mesh.vertices.len(), self.mesh.edges.len(), self.mesh.faces.len(), ) } /// 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 } } /// Run the demo. pub fn run_demo() -> WeatherForecast { let mesh_config = MeshConfig::from_refinement(4); // Smaller for demo let pred_config = PredictionConfig { num_steps: 10, time_step: 6.0, ..PredictionConfig::short_range() }; let mut weathercast = WeatherCast::new(mesh_config, pred_config); // Create sample initial state let initial_state = weathercast_shared::sample_atmospheric_state(); // Generate forecast weathercast.forecast(&initial_state) } // ============================================================================ // Tests // ============================================================================ #[cfg(test)] mod tests { use super::*; #[test] fn test_weathercast_creation() { let weathercast = WeatherCast::default_config(); assert!(!weathercast.trained); } #[test] fn test_forecast_generation() { let mut weathercast = WeatherCast::new( MeshConfig::from_refinement(3), PredictionConfig { num_steps: 5, ..Default::default() }, ); let initial = weathercast_shared::sample_atmospheric_state(); let forecast = weathercast.forecast(&initial); assert_eq!(forecast.forecast_hours(), 5); assert!(!forecast.hourly_states.is_empty()); } #[test] fn test_ensemble_forecast() { let mut weathercast = WeatherCast::new( MeshConfig::from_refinement(3), PredictionConfig { num_steps: 3, ..Default::default() }, ); let initial = weathercast_shared::sample_atmospheric_state(); let ensemble = weathercast.ensemble_forecast(&initial, 5); assert_eq!(ensemble.len(), 5); for (idx, forecast) in ensemble.iter().enumerate() { assert_eq!(forecast.ensemble_member, Some(idx)); } } #[test] fn test_verification() { let mut weathercast = WeatherCast::new( MeshConfig::from_refinement(3), PredictionConfig { num_steps: 3, ..Default::default() }, ); let initial = weathercast_shared::sample_atmospheric_state(); let forecast = weathercast.forecast(&initial); // Use initial state as "observations" for testing let observations = vec![initial.clone(), initial.clone(), initial]; let metrics = weathercast.verify(&forecast, &observations); assert!(!metrics.is_empty()); } #[test] fn test_uncertainty_calculation() { let mut weathercast = WeatherCast::new( MeshConfig::from_refinement(3), PredictionConfig { num_steps: 3, ..Default::default() }, ); let initial = weathercast_shared::sample_atmospheric_state(); let ensemble = weathercast.ensemble_forecast(&initial, 10); let uncertainty = weathercast.calculate_uncertainty(&ensemble, 0, "temperature"); assert!(!uncertainty.is_empty()); } #[test] fn test_training() { let mut weathercast = WeatherCast::new(MeshConfig::from_refinement(3), PredictionConfig::default()); let training_config = TrainingConfig { epochs: 3, ..Default::default() }; weathercast.train(&training_config, None); assert!(weathercast.is_trained()); } #[test] fn test_mesh_stats() { let weathercast = WeatherCast::new(MeshConfig::from_refinement(3), PredictionConfig::default()); let (nodes, edges, faces) = weathercast.mesh_stats(); assert!(nodes > 0); assert!(edges > 0); assert!(faces > 0); } #[test] fn test_run_demo() { let forecast = run_demo(); assert!(!forecast.hourly_states.is_empty()); assert!(!forecast.lead_times.is_empty()); } }