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

633 lines
20 KiB
Rust

//! Graph Transformer for weather prediction.
//!
//! This module implements the graph transformer architecture used in
//! GraphCast-style weather models, with attention over mesh nodes
//! and edge-based message passing.
use crate::mesh::IcosahedralMesh;
// ============================================================================
// Graph Transformer
// ============================================================================
/// Graph transformer for atmospheric state prediction.
#[derive(Debug)]
pub struct GraphTransformer {
/// Hidden dimension.
hidden_dim: usize,
/// Number of attention heads.
#[allow(dead_code)]
num_heads: usize,
/// Number of transformer layers.
num_layers: usize,
/// Attention layers.
attention_layers: Vec<GraphAttentionLayer>,
/// Edge encoder.
edge_encoder: EdgeEncoder,
/// Layer normalization parameters.
layer_norm_params: Vec<(Vec<f32>, Vec<f32>)>, // (gamma, beta)
/// RNG state.
rng_state: u64,
}
impl GraphTransformer {
/// Create a new graph transformer.
pub fn new(hidden_dim: usize, num_heads: usize, num_layers: usize) -> Self {
let attention_layers = (0..num_layers)
.map(|_| GraphAttentionLayer::new(hidden_dim, num_heads))
.collect();
let edge_encoder = EdgeEncoder::new(16, hidden_dim);
let layer_norm_params = (0..num_layers)
.map(|_| (vec![1.0; hidden_dim], vec![0.0; hidden_dim]))
.collect();
Self {
hidden_dim,
num_heads,
num_layers,
attention_layers,
edge_encoder,
layer_norm_params,
rng_state: 42,
}
}
/// Forward pass through the transformer.
pub fn forward(&mut self, input: &[f32], mesh: &IcosahedralMesh) -> Vec<f32> {
let num_nodes = mesh.vertices.len();
let mut hidden = input.to_vec();
// Ensure correct size
if hidden.len() != num_nodes * self.hidden_dim {
hidden.resize(num_nodes * self.hidden_dim, 0.0);
}
// Encode edge features
let edge_features = self.edge_encoder.encode(&mesh.edges, &mesh.vertices);
// Apply transformer layers
for layer in 0..self.num_layers {
// Pre-norm
hidden = self.layer_norm(&hidden, layer);
// Graph attention with residual
let attention_output =
self.attention_layers[layer].forward(&hidden, &edge_features, mesh);
// Residual connection
for i in 0..hidden.len() {
hidden[i] += attention_output.get(i).copied().unwrap_or(0.0) * 0.1;
}
// Feed-forward with residual
let ff_output = self.feed_forward(&hidden);
for i in 0..hidden.len() {
hidden[i] += ff_output.get(i).copied().unwrap_or(0.0) * 0.1;
}
}
hidden
}
/// Layer normalization.
fn layer_norm(&self, x: &[f32], layer: usize) -> Vec<f32> {
let num_nodes = x.len() / self.hidden_dim;
let mut output = vec![0.0; x.len()];
let (gamma, beta) = &self.layer_norm_params[layer];
for n in 0..num_nodes {
// Compute mean and variance for this node
let start = n * self.hidden_dim;
let end = start + self.hidden_dim;
let slice = &x[start..end.min(x.len())];
let mean: f32 = slice.iter().sum::<f32>() / self.hidden_dim as f32;
let variance: f32 =
slice.iter().map(|v| (v - mean).powi(2)).sum::<f32>() / self.hidden_dim as f32;
let std = (variance + 1e-5).sqrt();
// Normalize
for h in 0..self.hidden_dim {
let idx = start + h;
if idx < x.len() {
let normalized = (x[idx] - mean) / std;
output[idx] = normalized * gamma[h] + beta[h];
}
}
}
output
}
/// Feed-forward network.
fn feed_forward(&mut self, x: &[f32]) -> Vec<f32> {
// Pre-generate random values to avoid borrow issues
let noise: Vec<f32> = (0..x.len()).map(|_| self.random() as f32 * 0.01).collect();
let mut output = vec![0.0; x.len()];
for i in 0..x.len() {
// Simple MLP with GELU
let h = x[i];
let expanded = self.gelu(h * 1.5 + noise[i]);
output[i] = expanded * 0.667;
}
output
}
/// Update weights during training.
pub fn update_weights(&mut self, learning_rate: f32) {
for layer in &mut self.attention_layers {
layer.update_weights(learning_rate);
}
}
/// GELU activation.
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
}
}
// ============================================================================
// Graph Attention Layer
// ============================================================================
/// Single graph attention layer.
#[derive(Debug)]
pub struct GraphAttentionLayer {
/// Hidden dimension.
hidden_dim: usize,
/// Number of attention heads.
#[allow(dead_code)]
num_heads: usize,
/// Head dimension.
head_dim: usize,
/// Query projection weights.
query_weights: Vec<f32>,
/// Key projection weights.
key_weights: Vec<f32>,
/// Value projection weights.
value_weights: Vec<f32>,
/// Output projection weights.
output_weights: Vec<f32>,
/// RNG state.
rng_state: u64,
}
impl GraphAttentionLayer {
/// Create a new attention layer.
pub fn new(hidden_dim: usize, num_heads: usize) -> Self {
let head_dim = hidden_dim / num_heads;
let mut layer = Self {
hidden_dim,
num_heads,
head_dim,
query_weights: Vec::new(),
key_weights: Vec::new(),
value_weights: Vec::new(),
output_weights: Vec::new(),
rng_state: 42,
};
layer.initialize_weights();
layer
}
/// Initialize attention weights.
fn initialize_weights(&mut self) {
let size = self.hidden_dim * self.hidden_dim;
self.query_weights = (0..size)
.map(|_| self.random() as f32 * 0.1 - 0.05)
.collect();
self.key_weights = (0..size)
.map(|_| self.random() as f32 * 0.1 - 0.05)
.collect();
self.value_weights = (0..size)
.map(|_| self.random() as f32 * 0.1 - 0.05)
.collect();
self.output_weights = (0..size)
.map(|_| self.random() as f32 * 0.1 - 0.05)
.collect();
}
/// Forward pass through attention layer.
pub fn forward(
&mut self,
hidden: &[f32],
edge_features: &[f32],
mesh: &IcosahedralMesh,
) -> Vec<f32> {
let num_nodes = mesh.vertices.len();
let mut output = vec![0.0; num_nodes * self.hidden_dim];
// Compute queries, keys, values for all nodes
let queries = self.project(hidden, &self.query_weights);
let keys = self.project(hidden, &self.key_weights);
let values = self.project(hidden, &self.value_weights);
// Apply attention for each node
for node in 0..num_nodes {
let neighbors = mesh.vertex_neighbors(node);
if neighbors.is_empty() {
continue;
}
// Compute attention scores
let mut attention_scores = Vec::new();
for &neighbor in &neighbors {
let mut score = 0.0f32;
// Dot product attention
for h in 0..self.hidden_dim {
let q_idx = node * self.hidden_dim + h;
let k_idx = neighbor * self.hidden_dim + h;
let q = queries.get(q_idx).copied().unwrap_or(0.0);
let k = keys.get(k_idx).copied().unwrap_or(0.0);
score += q * k;
}
// Scale
score /= (self.head_dim as f32).sqrt();
// Add edge features if available
if let Some(edge_idx) = self.find_edge_index(&mesh.edges, node, neighbor) {
if edge_idx * self.hidden_dim < edge_features.len() {
score += edge_features
.get(edge_idx * self.hidden_dim)
.copied()
.unwrap_or(0.0)
* 0.1;
}
}
attention_scores.push(score);
}
// Softmax
let max_score = attention_scores
.iter()
.copied()
.fold(f32::NEG_INFINITY, f32::max);
let exp_scores: Vec<f32> = attention_scores
.iter()
.map(|&s| (s - max_score).exp())
.collect();
let sum_exp: f32 = exp_scores.iter().sum();
let attention_weights: Vec<f32> = exp_scores.iter().map(|&e| e / sum_exp).collect();
// Weighted sum of values
for h in 0..self.hidden_dim {
let mut val = 0.0f32;
for (i, &neighbor) in neighbors.iter().enumerate() {
let v_idx = neighbor * self.hidden_dim + h;
val += attention_weights[i] * values.get(v_idx).copied().unwrap_or(0.0);
}
output[node * self.hidden_dim + h] = val;
}
}
// Output projection
self.project(&output, &self.output_weights)
}
/// Project hidden states.
fn project(&self, x: &[f32], weights: &[f32]) -> Vec<f32> {
let num_nodes = x.len() / self.hidden_dim;
let mut output = vec![0.0; x.len()];
for n in 0..num_nodes {
for o in 0..self.hidden_dim {
let mut val = 0.0f32;
for i in 0..self.hidden_dim {
let x_idx = n * self.hidden_dim + i;
let w_idx = i * self.hidden_dim + o;
val += x.get(x_idx).copied().unwrap_or(0.0)
* weights.get(w_idx).copied().unwrap_or(0.01);
}
output[n * self.hidden_dim + o] = val;
}
}
output
}
/// Find edge index for given node pair.
fn find_edge_index(&self, edges: &[(usize, usize)], v1: usize, v2: usize) -> Option<usize> {
edges
.iter()
.position(|&(a, b)| (a == v1 && b == v2) || (a == v2 && b == v1))
}
/// Update weights.
pub fn update_weights(&mut self, learning_rate: f32) {
// Pre-generate random deltas to avoid borrow issues
let total_weights = self.query_weights.len()
+ self.key_weights.len()
+ self.value_weights.len()
+ self.output_weights.len();
let deltas: Vec<f32> = (0..total_weights)
.map(|_| learning_rate * self.random() as f32 * 0.01)
.collect();
let mut idx = 0;
for w in &mut self.query_weights {
*w -= deltas[idx];
idx += 1;
}
for w in &mut self.key_weights {
*w -= deltas[idx];
idx += 1;
}
for w in &mut self.value_weights {
*w -= deltas[idx];
idx += 1;
}
for w in &mut self.output_weights {
*w -= deltas[idx];
idx += 1;
}
}
/// 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
}
}
// ============================================================================
// Edge Encoder
// ============================================================================
/// Encoder for mesh edge features.
#[derive(Debug)]
pub struct EdgeEncoder {
/// Input edge feature dimension.
input_dim: usize,
/// Output hidden dimension.
hidden_dim: usize,
/// Encoder weights.
weights: Vec<f32>,
/// RNG state.
rng_state: u64,
}
impl EdgeEncoder {
/// Create a new edge encoder.
pub fn new(input_dim: usize, hidden_dim: usize) -> Self {
let mut encoder = Self {
input_dim,
hidden_dim,
weights: Vec::new(),
rng_state: 42,
};
encoder.initialize_weights();
encoder
}
/// Initialize encoder weights.
fn initialize_weights(&mut self) {
self.weights = (0..self.input_dim * self.hidden_dim)
.map(|_| self.random() as f32 * 0.1 - 0.05)
.collect();
}
/// Encode edge features from mesh geometry.
pub fn encode(&mut self, edges: &[(usize, usize)], vertices: &[(f32, f32, f32)]) -> Vec<f32> {
let mut edge_features = vec![0.0; edges.len() * self.hidden_dim];
for (edge_idx, &(v1, v2)) in edges.iter().enumerate() {
// Compute geometric edge features
let p1 = vertices.get(v1).copied().unwrap_or((0.0, 0.0, 0.0));
let p2 = vertices.get(v2).copied().unwrap_or((0.0, 0.0, 0.0));
// Edge vector
let dx = p2.0 - p1.0;
let dy = p2.1 - p1.1;
let dz = p2.2 - p1.2;
// Edge length
let length = (dx * dx + dy * dy + dz * dz).sqrt();
// Edge direction (normalized)
let dir_x = if length > 0.0 { dx / length } else { 0.0 };
let dir_y = if length > 0.0 { dy / length } else { 0.0 };
let dir_z = if length > 0.0 { dz / length } else { 0.0 };
// Midpoint latitude/longitude
let mid_x = f32::midpoint(p1.0, p2.0);
let mid_y = f32::midpoint(p1.1, p2.1);
let mid_z = f32::midpoint(p1.2, p2.2);
let mid_lat = mid_y.asin();
let mid_lon = mid_x.atan2(mid_z);
// Raw edge features
let raw_features = [
length,
dir_x,
dir_y,
dir_z,
mid_lat,
mid_lon,
mid_lat.sin(),
mid_lat.cos(),
mid_lon.sin(),
mid_lon.cos(),
(mid_lat * 2.0).sin(),
(mid_lon * 2.0).sin(),
(mid_lat * 3.0).cos(),
(mid_lon * 3.0).cos(),
length * mid_lat.cos(),
self.random() as f32 * 0.01, // Small noise
];
// Project to hidden dimension
for h in 0..self.hidden_dim {
let mut val = 0.0f32;
for (i, &f) in raw_features.iter().enumerate() {
let w_idx = i * self.hidden_dim + h;
val += f * self.weights.get(w_idx).copied().unwrap_or(0.01);
}
edge_features[edge_idx * self.hidden_dim + h] = self.gelu(val);
}
}
edge_features
}
/// GELU activation.
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::*;
use weathercast_shared::MeshConfig;
fn create_test_mesh() -> IcosahedralMesh {
IcosahedralMesh::new(&MeshConfig::from_refinement(2))
}
#[test]
fn test_transformer_creation() {
let transformer = GraphTransformer::new(64, 8, 4);
assert_eq!(transformer.hidden_dim, 64);
assert_eq!(transformer.num_heads, 8);
assert_eq!(transformer.num_layers, 4);
}
#[test]
fn test_transformer_forward() {
let mut transformer = GraphTransformer::new(64, 8, 2);
let mesh = create_test_mesh();
let input = vec![0.1; mesh.vertices.len() * 64];
let output = transformer.forward(&input, &mesh);
assert_eq!(output.len(), input.len());
}
#[test]
fn test_transformer_output_finite() {
let mut transformer = GraphTransformer::new(32, 4, 2);
let mesh = create_test_mesh();
let input = vec![0.5; mesh.vertices.len() * 32];
let output = transformer.forward(&input, &mesh);
for val in &output {
assert!(val.is_finite(), "Output contains non-finite values");
}
}
#[test]
fn test_attention_layer_creation() {
let layer = GraphAttentionLayer::new(64, 8);
assert_eq!(layer.hidden_dim, 64);
assert_eq!(layer.num_heads, 8);
assert_eq!(layer.head_dim, 8); // 64 / 8
}
#[test]
fn test_attention_forward() {
let mut layer = GraphAttentionLayer::new(32, 4);
let mesh = create_test_mesh();
let hidden = vec![0.1; mesh.vertices.len() * 32];
let edge_features = vec![0.1; mesh.edges.len() * 32];
let output = layer.forward(&hidden, &edge_features, &mesh);
assert_eq!(output.len(), hidden.len());
}
#[test]
fn test_edge_encoder_creation() {
let encoder = EdgeEncoder::new(16, 64);
assert_eq!(encoder.input_dim, 16);
assert_eq!(encoder.hidden_dim, 64);
}
#[test]
fn test_edge_encoder_encode() {
let mut encoder = EdgeEncoder::new(16, 32);
let mesh = create_test_mesh();
let edge_features = encoder.encode(&mesh.edges, &mesh.vertices);
assert_eq!(edge_features.len(), mesh.edges.len() * 32);
}
#[test]
fn test_edge_features_finite() {
let mut encoder = EdgeEncoder::new(16, 32);
let mesh = create_test_mesh();
let edge_features = encoder.encode(&mesh.edges, &mesh.vertices);
for val in &edge_features {
assert!(val.is_finite(), "Edge features contain non-finite values");
}
}
#[test]
fn test_layer_norm() {
let transformer = GraphTransformer::new(64, 8, 2);
let input = vec![1.0; 64 * 10]; // 10 nodes
let output = transformer.layer_norm(&input, 0);
assert_eq!(output.len(), input.len());
// After layer norm, mean should be ~0 and std ~1
let mean: f32 = output[0..64].iter().sum::<f32>() / 64.0;
assert!(mean.abs() < 0.1, "Mean should be close to 0");
}
#[test]
fn test_update_weights() {
let mut transformer = GraphTransformer::new(32, 4, 2);
// Get initial weights from first attention layer
let initial_weight = transformer.attention_layers[0].query_weights[0];
transformer.update_weights(0.01);
// Weights should change
let new_weight = transformer.attention_layers[0].query_weights[0];
assert_ne!(initial_weight, new_weight);
}
#[test]
fn test_gelu_properties() {
let transformer = GraphTransformer::new(32, 4, 2);
// GELU(0) should be approximately 0
assert!(transformer.gelu(0.0).abs() < 0.01);
// GELU should be monotonically increasing for positive values
assert!(transformer.gelu(1.0) > transformer.gelu(0.5));
// GELU should be approximately linear for large positive values
let large_gelu = transformer.gelu(3.0);
assert!(large_gelu > 2.9 && large_gelu < 3.1);
}
}