Files
rustytorch/demos/rtx-weathercast-demo/src/mesh.rs
T
2026-03-04 00:08:42 +00:00

565 lines
19 KiB
Rust

//! 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<usize> {
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<f32>,
/// Decoder weights.
decoder_weights: Vec<f32>,
/// 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<f32> {
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<f32> {
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));
}
}