Initial commit
This commit is contained in:
@@ -0,0 +1,632 @@
|
||||
//! 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);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user