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