Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
+427
View File
@@ -0,0 +1,427 @@
//! 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(&current_state, &self.mesh);
// Decode mesh state back to atmospheric state
let atm_state = self
.mesh_encoder
.decode_state(&current_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());
}
}