//! Physics-Informed Neural Operator (PINO) for cardiac electrophysiology. //! //! Implements a neural operator that learns the solution operator for the //! monodomain equation with physics-based constraints. use cardiosim_shared::{ HeartMesh, NeuralOperatorConfig, SimulationConfig, StimulationProtocol, VoltageField, }; use crate::CardioSimError; /// Physics-Informed Neural Operator for cardiac simulation. #[derive(Debug)] pub struct CardiacPINO { config: NeuralOperatorConfig, fno_layers: Vec, lifting: MLPLayer, projection: MLPLayer, } impl CardiacPINO { /// Create a new cardiac PINO. #[must_use] pub fn new(config: NeuralOperatorConfig) -> Self { let lifting = MLPLayer::new(config.input_channels, config.hidden_dim); let fno_layers: Vec = (0..config.num_layers) .map(|_| FourierLayer::new(config.hidden_dim, 16)) .collect(); let projection = MLPLayer::new(config.hidden_dim, config.output_channels); Self { config, fno_layers, lifting, projection, } } /// Solve using neural operator. pub fn solve( &self, mesh: &HeartMesh, protocol: &StimulationProtocol, config: &SimulationConfig, ) -> Result, CardioSimError> { let n_vertices = mesh.vertices.len(); let n_outputs = (config.total_time / config.output_interval) as usize; // Initialize input (voltage, recovery, coordinates) let mut voltage = vec![-85.0_f32; n_vertices]; let recovery = vec![0.0_f32; n_vertices]; // Encode mesh coordinates let coords: Vec<[f32; 3]> = mesh .vertices .iter() .map(|v| [v.x / 100.0, v.y / 100.0, v.z / 100.0]) // Normalize .collect(); let mut output = Vec::new(); // Simulate in chunks using neural operator let _time_steps_per_inference = self.config.time_steps; let mut t = 0.0; for _ in 0..n_outputs { // Prepare input tensor let input = self.prepare_input(&voltage, &recovery, &coords, mesh, protocol, t); // Forward through neural operator let prediction = self.forward(&input); // Extract voltage and recovery let new_voltage: Vec = prediction.iter().take(n_vertices).copied().collect(); let _new_recovery: Vec = prediction .iter() .skip(n_vertices) .take(n_vertices) .copied() .collect(); voltage = new_voltage; output.push(VoltageField { time: t, voltages: voltage.clone(), }); t += config.output_interval; } Ok(output) } fn prepare_input( &self, voltage: &[f32], recovery: &[f32], coords: &[[f32; 3]], mesh: &HeartMesh, protocol: &StimulationProtocol, t: f32, ) -> Vec { let n = voltage.len(); let mut input = Vec::with_capacity(n * self.config.input_channels); for i in 0..n { // Voltage (normalized) input.push((voltage[i] + 85.0) / 120.0); // Recovery input.push(recovery[i]); // Spatial coordinates input.push(coords[i][0]); input.push(coords[i][1]); input.push(coords[i][2]); // Fiber direction input.push(mesh.fibers[i].x); input.push(mesh.fibers[i].y); input.push(mesh.fibers[i].z); } // Add stimulus encoding for site in &protocol.sites { for &stim_time in &site.times { if t >= stim_time && t < stim_time + site.duration { // Mark stimulated vertices for (i, vertex) in mesh.vertices.iter().enumerate() { let dist = vertex.distance_to(&site.center); if dist < site.radius { input[i * self.config.input_channels] += 0.5; } } } } } input } fn forward(&self, input: &[f32]) -> Vec { // Lifting layer let mut hidden = self.lifting.forward(input); // Fourier layers for layer in &self.fno_layers { hidden = layer.forward(&hidden); } // Projection layer self.projection.forward(&hidden) } } /// Fourier layer for FNO. #[derive(Debug)] struct FourierLayer { hidden_dim: usize, modes: usize, weights_real: Vec>, weights_imag: Vec>, linear_weights: Vec>, } impl FourierLayer { fn new(hidden_dim: usize, modes: usize) -> Self { use rand::SeedableRng; use rand_distr::{Distribution, Normal}; let mut rng = rand::rngs::StdRng::seed_from_u64(42); let std = (2.0 / hidden_dim as f32).sqrt(); let normal = Normal::new(0.0_f32, std).unwrap(); let weights_real: Vec> = (0..modes) .map(|_| (0..hidden_dim).map(|_| normal.sample(&mut rng)).collect()) .collect(); let weights_imag: Vec> = (0..modes) .map(|_| (0..hidden_dim).map(|_| normal.sample(&mut rng)).collect()) .collect(); let linear_weights: Vec> = (0..hidden_dim) .map(|_| (0..hidden_dim).map(|_| normal.sample(&mut rng)).collect()) .collect(); Self { hidden_dim, modes, weights_real, weights_imag, linear_weights, } } fn forward(&self, input: &[f32]) -> Vec { let n = input.len() / self.hidden_dim; if n == 0 { return input.to_vec(); } // Simplified spectral convolution let mut output = vec![0.0_f32; input.len()]; // Linear path (skip connection) for i in 0..n { for j in 0..self.hidden_dim { let input_val = input.get(i * self.hidden_dim + j).copied().unwrap_or(0.0); for k in 0..self.hidden_dim { let weight = self .linear_weights .get(j) .and_then(|w| w.get(k)) .copied() .unwrap_or(0.0); output[i * self.hidden_dim + k] += input_val * weight; } } } // Spectral path (simplified DFT-like operation) for mode in 0..self.modes.min(n) { let freq = 2.0 * std::f32::consts::PI * mode as f32 / n as f32; for i in 0..n { let phase = freq * i as f32; let cos_p = phase.cos(); let sin_p = phase.sin(); for j in 0..self.hidden_dim { let input_val = input.get(i * self.hidden_dim + j).copied().unwrap_or(0.0); let weight_r = self .weights_real .get(mode) .and_then(|w| w.get(j)) .copied() .unwrap_or(0.0); let weight_i = self .weights_imag .get(mode) .and_then(|w| w.get(j)) .copied() .unwrap_or(0.0); output[i * self.hidden_dim + j] += input_val * (weight_r * cos_p - weight_i * sin_p); } } } // GELU activation for v in &mut output { *v = gelu(*v); } output } } /// MLP layer. #[derive(Debug)] struct MLPLayer { in_dim: usize, out_dim: usize, weights: Vec>, bias: Vec, } impl MLPLayer { fn new(in_dim: usize, out_dim: usize) -> Self { use rand::SeedableRng; use rand_distr::{Distribution, Normal}; let mut rng = rand::rngs::StdRng::seed_from_u64(123); let std = (2.0 / in_dim as f32).sqrt(); let normal = Normal::new(0.0_f32, std).unwrap(); let weights: Vec> = (0..in_dim) .map(|_| (0..out_dim).map(|_| normal.sample(&mut rng)).collect()) .collect(); let bias: Vec = (0..out_dim).map(|_| 0.0).collect(); Self { in_dim, out_dim, weights, bias, } } fn forward(&self, input: &[f32]) -> Vec { let batch_size = input.len() / self.in_dim; if batch_size == 0 { return vec![0.0; self.out_dim]; } let mut output = vec![0.0_f32; batch_size * self.out_dim]; for b in 0..batch_size { for j in 0..self.out_dim { let mut sum = self.bias[j]; for i in 0..self.in_dim { let input_idx = b * self.in_dim + i; let input_val = input.get(input_idx).copied().unwrap_or(0.0); let weight = self .weights .get(i) .and_then(|w| w.get(j)) .copied() .unwrap_or(0.0); sum += input_val * weight; } output[b * self.out_dim + j] = sum; } } output } } /// GELU activation function. fn gelu(x: f32) -> f32 { 0.5 * x * (1.0 + ((2.0_f32 / std::f32::consts::PI).sqrt() * (x + 0.044715 * x.powi(3))).tanh()) } #[cfg(test)] mod tests { use super::*; #[test] fn test_pino_creation() { let config = NeuralOperatorConfig::default(); let pino = CardiacPINO::new(config); assert_eq!(pino.fno_layers.len(), 4); } #[test] fn test_pino_solve() { let config = NeuralOperatorConfig::default(); let pino = CardiacPINO::new(config); let mesh = cardiosim_shared::get_sample_heart_mesh(); let protocol = cardiosim_shared::get_sample_protocol(); let sim_config = SimulationConfig { total_time: 10.0, output_interval: 5.0, ..Default::default() }; let result = pino.solve(&mesh, &protocol, &sim_config); assert!(result.is_ok()); let fields = result.unwrap(); assert!(!fields.is_empty()); } #[test] fn test_fourier_layer() { let layer = FourierLayer::new(32, 8); let input = vec![0.1_f32; 32 * 10]; let output = layer.forward(&input); assert_eq!(output.len(), input.len()); } #[test] fn test_mlp_layer() { let layer = MLPLayer::new(16, 32); let input = vec![0.1_f32; 16 * 5]; let output = layer.forward(&input); assert_eq!(output.len(), 32 * 5); } #[test] fn test_gelu() { assert!((gelu(0.0) - 0.0).abs() < 0.01); assert!(gelu(2.0) > 1.9); assert!(gelu(-2.0) < 0.1); } }