428 lines
13 KiB
Rust
428 lines
13 KiB
Rust
//! 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<GridPoint> = 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<WeatherForecast> {
|
|
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<VerificationMetrics> {
|
|
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<UncertaintyEstimate> {
|
|
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<f32> = 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<Box<dyn Fn(TrainingProgress) + Send>>,
|
|
) {
|
|
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());
|
|
}
|
|
}
|