//! Core PINN (Physics-Informed Neural Network) Implementation //! //! This module provides the main PINN structure and functionality for solving //! partial differential equations with neural networks that respect physics laws. use crate::error::{Result, ScienceError}; use crate::physics::{ BoundaryConditions, ConservationLoss, PINNTrainer, PhysicsLoss, TrainingConfig, }; use crate::variable_extensions::VariableExt; use async_trait::async_trait; use rtx_autograd::Variable; use rtx_tensor::{Device, Tensor, TensorError}; use serde::{Deserialize, Serialize}; use std::collections::HashMap; use tokio::sync::RwLock; use tracing::{debug, info, warn}; /// Physics-Informed Neural Network for solving PDEs pub struct PINN { /// Neural network layers network: Vec, /// Device for computations device: Device, /// Physics loss function physics_loss: Box, /// Conservation losses conservation_losses: Vec>, /// Training configuration config: TrainingConfig, /// Current training state state: RwLock, } /// Individual linear layer in the neural network #[derive(Debug, Clone)] struct LinearLayer { /// Weight matrix weights: Variable, /// Bias vector bias: Variable, /// Activation function activation: ActivationType, } /// Supported activation functions #[derive(Debug, Clone, Serialize, Deserialize)] pub enum ActivationType { /// Hyperbolic tangent Tanh, /// Sigmoid Sigmoid, /// Sine activation (useful for periodic solutions) Sin, /// Swish/SiLU activation Swish, /// GELU activation GELU, /// `ReLU` activation ReLU, /// Learnable activation (adaptive) Learnable, } /// Internal state of the PINN during training #[derive(Debug, Default, Clone)] pub struct PINNState { /// Current training iteration iteration: usize, /// Training loss history loss_history: Vec, /// Best model weights (for early stopping) best_weights: Option>, /// Best validation loss best_loss: f64, /// Convergence metrics convergence: ConvergenceMetrics, } /// Components of the total loss function #[derive(Debug, Clone)] pub struct LossComponents { /// Physics/PDE loss pub physics: f64, /// Boundary condition loss pub boundary: f64, /// Initial condition loss pub initial: f64, /// Conservation law losses pub conservation: Vec, /// Data fitting loss (if available) pub data: f64, /// Total combined loss pub total: f64, /// Training iteration pub iteration: usize, } /// Convergence tracking metrics #[derive(Debug, Default, Clone)] struct ConvergenceMetrics { /// Moving average of loss gradient magnitude gradient_norm: f64, /// Loss improvement rate improvement_rate: f64, /// Plateau detection counter plateau_count: usize, /// Early stopping criteria met should_stop: bool, } /// Checkpoint for saving and restoring PINN state #[derive(Debug, Clone)] pub struct PINNCheckpoint { /// Layer weights pub weights: Vec, /// Layer biases pub biases: Vec, /// Training iteration pub iteration: usize, /// Best loss value pub best_loss: f64, } /// Builder pattern for PINN construction pub struct PINNBuilder { device: Option, layers: Vec, activation: ActivationType, physics_loss: Option>, conservation_losses: Vec>, config: TrainingConfig, } /// Trait for physics-informed neural network functionality #[async_trait] pub trait PhysicsInformedNetwork: Send + Sync { /// Predict output for given input coordinates async fn predict(&self, inputs: &[(f64, f64)]) -> Result>; /// Predict with gradients (for physics loss computation) async fn predict_with_gradients( &self, inputs: &[(f64, f64)], ) -> Result<(Vec, Vec, Vec)>; /// Train the network with boundary conditions async fn train(&mut self, boundary_data: BoundaryConditions, epochs: usize) -> Result<()>; /// Evaluate physics loss at given points async fn physics_residual(&self, inputs: &[(f64, f64)]) -> Result>; /// Get current training state async fn training_state(&self) -> Result; } impl PINN { /// Create a new PINN builder #[must_use] pub fn builder() -> PINNBuilder { PINNBuilder::new() } /// Create PINN from builder pub(crate) fn from_builder(builder: PINNBuilder) -> Result { let device = builder.device.ok_or_else(|| { ScienceError::physics( "Device must be specified", crate::error::PhysicsDomain::FluidDynamics, ) })?; let physics_loss = builder.physics_loss.ok_or_else(|| { ScienceError::physics( "Physics loss must be specified", crate::error::PhysicsDomain::FluidDynamics, ) })?; if builder.layers.len() < 2 { return Err(ScienceError::physics( "At least input and output layers required", crate::error::PhysicsDomain::FluidDynamics, )); } // Initialize network layers let mut network = Vec::new(); for i in 0..builder.layers.len() - 1 { let input_size = builder.layers[i]; let output_size = builder.layers[i + 1]; let layer = LinearLayer::new(input_size, output_size, &builder.activation, &device)?; network.push(layer); } Ok(Self { network, device, physics_loss, conservation_losses: builder.conservation_losses, config: builder.config, state: RwLock::new(PINNState::default()), }) } /// Forward pass through the network pub async fn forward(&self, inputs: &Tensor) -> Result { let mut x = Variable::from_tensor(inputs.clone()); for (i, layer) in self.network.iter().enumerate() { x = layer.forward(x)?; // Apply activation (skip for output layer) if i < self.network.len() - 1 { x = apply_activation(x, &layer.activation)?; } } Ok(x) } /// Compute derivatives using automatic differentiation async fn compute_derivatives( &self, inputs: &Tensor, outputs: &Variable, ) -> Result<(Variable, Variable)> { // Simplified gradient computation using finite differences // A full implementation would use automatic differentiation let epsilon = 1e-5; // Create perturbation tensors let eps_x = Tensor::from_slice(&[epsilon, 0.0], &[1, 2], &self.device)?; let eps_t = Tensor::from_slice(&[0.0, epsilon], &[1, 2], &self.device)?; // Perturb inputs let x_plus = inputs.add(&eps_x)?; let t_plus = inputs.add(&eps_t)?; // Forward pass with perturbed inputs let output_x_plus = self.forward(&x_plus).await?; let output_t_plus = self.forward(&t_plus).await?; // Compute finite difference approximations let du_dx_tensor = output_x_plus .tensor() .sub(outputs.tensor())? .mul_scalar(1.0 / epsilon)?; let du_dt_tensor = output_t_plus .tensor() .sub(outputs.tensor())? .mul_scalar(1.0 / epsilon)?; // Create Variables from the gradient tensors let du_dx = Variable::new(du_dx_tensor, false); let du_dt = Variable::new(du_dt_tensor, false); Ok((du_dx, du_dt)) } /// Compute second derivatives async fn compute_second_derivatives( &self, inputs: &Tensor, first_derivs: &(Variable, Variable), ) -> Result<(Variable, Variable, Variable)> { // Simplified second derivative computation using finite differences let epsilon = 1e-5; let (du_dx, du_dt) = first_derivs; // Create perturbation tensors let eps_x = Tensor::from_slice(&[epsilon, 0.0], &[1, 2], &self.device)?; let eps_t = Tensor::from_slice(&[0.0, epsilon], &[1, 2], &self.device)?; // Perturb inputs in both directions let x_plus = inputs.add(&eps_x)?; let x_minus = inputs.sub(&eps_x)?; let t_plus = inputs.add(&eps_t)?; let t_minus = inputs.sub(&eps_t)?; // Compute outputs at perturbed points let out_x_plus = self.forward(&x_plus).await?; let out_x_minus = self.forward(&x_minus).await?; let out_t_plus = self.forward(&t_plus).await?; let out_t_minus = self.forward(&t_minus).await?; // Compute second derivatives using central differences let two_du_dx = du_dx.tensor().mul_scalar(2.0)?; let two_du_dt = du_dt.tensor().mul_scalar(2.0)?; let d2u_dx2_tensor = out_x_plus .tensor() .add(out_x_minus.tensor())? .sub(&two_du_dx)? .mul_scalar(1.0 / (epsilon * epsilon))?; let d2u_dt2_tensor = out_t_plus .tensor() .add(out_t_minus.tensor())? .sub(&two_du_dt)? .mul_scalar(1.0 / (epsilon * epsilon))?; // Mixed derivative (simplified) let d2u_dxdt_tensor = out_x_plus .tensor() .sub(out_x_minus.tensor())? .mul_scalar(1.0 / (2.0 * epsilon * epsilon))?; // Create Variables from the gradient tensors let d2u_dx2 = Variable::new(d2u_dx2_tensor, false); let d2u_dt2 = Variable::new(d2u_dt2_tensor, false); let d2u_dxdt = Variable::new(d2u_dxdt_tensor, false); Ok((d2u_dx2, d2u_dt2, d2u_dxdt)) } /// Compute total loss combining all components async fn compute_total_loss( &self, collocation_points: &Tensor, boundary_data: &BoundaryConditions, ) -> Result<(Variable, LossComponents)> { let mut loss_components = LossComponents { physics: 0.0, boundary: 0.0, initial: 0.0, conservation: Vec::new(), data: 0.0, total: 0.0, iteration: 0, }; // Physics/PDE loss let outputs = self.forward(collocation_points).await?; let (du_dx, du_dt) = self .compute_derivatives(collocation_points, &outputs) .await?; let second_derivs = self .compute_second_derivatives(collocation_points, &(du_dx.clone(), du_dt.clone())) .await?; let physics_residual = self .physics_loss .compute_residual( collocation_points, &outputs, &du_dx, &du_dt, &second_derivs.0, &second_derivs.1, &second_derivs.2, ) .await?; // Use the mean_square from rtx_autograd Variable let squared = physics_residual.multiply(&physics_residual)?; let physics_loss_var = squared.mean()?; let physics_loss = physics_loss_var.value().to_scalar::()?; loss_components.physics = f64::from(physics_loss); // Boundary condition loss let boundary_loss_var = boundary_data.compute_loss(self).await?; let boundary_loss = boundary_loss_var.value().to_scalar::()?; loss_components.boundary = f64::from(boundary_loss); // Conservation law losses for conservation_law in &self.conservation_losses { let conservation_loss = conservation_law .compute_loss(collocation_points, &outputs, &du_dx, &du_dt) .await?; loss_components .conservation .push(f64::from(conservation_loss)); } // Combine losses with weights let mut total_loss = physics_loss * (self.config.physics_weight as f32); total_loss += boundary_loss * self.config.boundary_weight as f32; for (i, conservation_law) in self.conservation_losses.iter().enumerate() { if let Some(weight) = self.config.conservation_weights.get(i) { let cons_loss = conservation_law .compute_loss(collocation_points, &outputs, &du_dx, &du_dt) .await?; total_loss += cons_loss * *weight as f32; } } loss_components.total = f64::from(total_loss); // Convert total_loss to Variable for backward pass let loss_tensor = Tensor::scalar(total_loss, rtx_tensor::DType::F32, &self.device)?; let loss_variable = Variable::from_tensor(loss_tensor); Ok((loss_variable, loss_components)) } /// Update training state with new loss information async fn update_training_state(&self, loss_components: LossComponents) -> Result<()> { let mut state = self.state.write().await; state.iteration += 1; // Update convergence metrics if let Some(prev_loss) = state.loss_history.last() { let loss_change = prev_loss.total - loss_components.total; state.convergence.improvement_rate = loss_change / prev_loss.total; if loss_change.abs() < self.config.convergence_tolerance { state.convergence.plateau_count += 1; } else { state.convergence.plateau_count = 0; } // Early stopping check if state.convergence.plateau_count > self.config.patience { state.convergence.should_stop = true; warn!("Early stopping triggered due to lack of improvement"); } } // Save best model if loss_components.total < state.best_loss { state.best_loss = loss_components.total; // Save current model weights as best weights let mut best_weights = HashMap::new(); for (i, layer) in self.network.iter().enumerate() { let weight_key = format!("layer_{i}_weight"); let bias_key = format!("layer_{i}_bias"); best_weights.insert(weight_key, layer.weights.value().clone()); best_weights.insert(bias_key, layer.bias.value().clone()); } state.best_weights = Some(best_weights); debug!( "New best model saved at iteration {} with loss {:.6}", state.iteration, loss_components.total ); } if state.iteration % 100 == 0 { info!( "Iteration {}: Total loss = {:.6}, Physics = {:.6}, Boundary = {:.6}", state.iteration, loss_components.total, loss_components.physics, loss_components.boundary ); } state.loss_history.push(loss_components); Ok(()) } } #[async_trait] impl PhysicsInformedNetwork for PINN { async fn predict(&self, inputs: &[(f64, f64)]) -> Result> { let input_tensor = create_input_tensor(inputs, &self.device)?; let output = self.forward(&input_tensor).await?; let values = output .tensor() .to_vec()? .iter() .map(|&x| f64::from(x)) .collect(); Ok(values) } async fn predict_with_gradients( &self, inputs: &[(f64, f64)], ) -> Result<(Vec, Vec, Vec)> { let input_tensor = create_input_tensor(inputs, &self.device)?; let output = self.forward(&input_tensor).await?; let (du_dx, du_dt) = self.compute_derivatives(&input_tensor, &output).await?; let values = output .tensor() .to_vec()? .iter() .map(|&x| f64::from(x)) .collect(); let dx_values = du_dx .tensor() .to_vec()? .iter() .map(|&x| f64::from(x)) .collect(); let dt_values = du_dt .tensor() .to_vec()? .iter() .map(|&x| f64::from(x)) .collect(); Ok((values, dx_values, dt_values)) } async fn train(&mut self, boundary_data: BoundaryConditions, epochs: usize) -> Result<()> { let trainer = PINNTrainer::new(self.config.clone()); trainer.train(self, boundary_data, epochs).await } async fn physics_residual(&self, inputs: &[(f64, f64)]) -> Result> { let input_tensor = create_input_tensor(inputs, &self.device)?; let output = self.forward(&input_tensor).await?; let (du_dx, du_dt) = self.compute_derivatives(&input_tensor, &output).await?; let derivatives = (du_dx, du_dt); let second_derivs = self .compute_second_derivatives(&input_tensor, &derivatives) .await?; let residual = self .physics_loss .compute_residual( &input_tensor, &output, &derivatives.0, &derivatives.1, &second_derivs.0, &second_derivs.1, &second_derivs.2, ) .await?; let values = residual .tensor() .to_vec()? .iter() .map(|&x| f64::from(x)) .collect(); Ok(values) } async fn training_state(&self) -> Result { let state_guard = self.state.read().await; Ok((*state_guard).clone()) } } impl PINN { /// Save model weights to a file using `SafeTensors` format pub async fn save_weights(&self, path: &std::path::Path) -> Result<()> { use std::collections::HashMap; use std::io::Write; let mut weights_map = HashMap::new(); // Collect all layer weights and biases for (i, layer) in self.network.iter().enumerate() { let weight_key = format!("layer_{i}_weight"); let bias_key = format!("layer_{i}_bias"); // Get tensor data from the layer let weight_data = layer.weights.value().to_vec()?; let bias_data = layer.bias.value().to_vec()?; weights_map.insert(weight_key, weight_data); weights_map.insert(bias_key, bias_data); } // For now, use a simple JSON serialization // In a production system, this would use SafeTensors format let json_data = serde_json::to_string_pretty(&weights_map)?; let mut file = std::fs::File::create(path)?; file.write_all(json_data.as_bytes())?; tracing::info!("Model weights saved to {:?}", path); Ok(()) } /// Load model weights from a file pub async fn load_weights(&mut self, path: &std::path::Path) -> Result<()> { use std::collections::HashMap; if !path.exists() { return Err(ScienceError::io_error( format!("Weight file not found: {path:?}"), "file_not_found", )); } let json_data = std::fs::read_to_string(path)?; let weights_map: HashMap> = serde_json::from_str(&json_data)?; // Load weights into each layer for (i, layer) in self.network.iter_mut().enumerate() { let weight_key = format!("layer_{i}_weight"); let bias_key = format!("layer_{i}_bias"); if let Some(weight_data) = weights_map.get(&weight_key) { // Create new tensor from the loaded data let weight_shape = layer.weights.value().shape().to_vec(); let weight_tensor = Tensor::from_slice(weight_data, &weight_shape, &self.device)?; layer.weights = Variable::new(weight_tensor, true); } if let Some(bias_data) = weights_map.get(&bias_key) { let bias_shape = layer.bias.value().shape().to_vec(); let bias_tensor = Tensor::from_slice(bias_data, &bias_shape, &self.device)?; layer.bias = Variable::new(bias_tensor, true); } } tracing::info!("Model weights loaded from {:?}", path); Ok(()) } /// Get device used by this PINN pub fn device(&self) -> &Device { &self.device } /// Forward pass with gradient tracking for physics-informed training pub async fn forward_with_gradients( &mut self, inputs: &Tensor, ) -> Result<(Variable, Variable, Variable)> { let outputs = self.forward(inputs).await?; let (du_dx, du_dt) = self.compute_derivatives(inputs, &outputs).await?; Ok((outputs, du_dx, du_dt)) } /// Get all trainable parameters pub fn parameters(&self) -> Vec { let mut params = Vec::new(); for layer in &self.network { params.push(layer.weights.clone()); params.push(layer.bias.clone()); } params } /// Zero all gradients pub fn zero_gradients(&mut self) { for layer in &mut self.network { layer.weights.zero_grad(); layer.bias.zero_grad(); } } /// Clip gradients to prevent explosion pub fn clip_gradients(&mut self, max_norm: f32) { for layer in &mut self.network { layer.weights.clip_grad(max_norm); layer.bias.clip_grad(max_norm); } } /// Save current state for checkpointing pub async fn save_state(&self) -> Result { let mut weights = Vec::new(); let mut biases = Vec::new(); for layer in &self.network { weights.push(layer.weights.value().clone()); biases.push(layer.bias.value().clone()); } let state = self.state.read().await; Ok(PINNCheckpoint { weights, biases, iteration: state.iteration, best_loss: state.best_loss, }) } /// Restore from checkpoint pub async fn restore_state(&mut self, checkpoint: PINNCheckpoint) -> Result<()> { if checkpoint.weights.len() != self.network.len() { return Err(ScienceError::physics( "Checkpoint layer count mismatch", crate::error::PhysicsDomain::FluidDynamics, )); } for (i, layer) in self.network.iter_mut().enumerate() { layer.weights = Variable::from_tensor(checkpoint.weights[i].clone()); layer.bias = Variable::from_tensor(checkpoint.biases[i].clone()); } let mut state = self.state.write().await; state.iteration = checkpoint.iteration; state.best_loss = checkpoint.best_loss; Ok(()) } /// Predict without gradients for inference pub async fn predict(&self, inputs: &Tensor) -> Result { let mut x = inputs.clone(); for (i, layer) in self.network.iter().enumerate() { // Manual forward pass without Variable tracking let weights_tensor = layer.weights.value(); let bias_tensor = layer.bias.value(); x = x.matmul(weights_tensor)?; x = x.add(bias_tensor)?; // Apply activation (skip for output layer) if i < self.network.len() - 1 { x = apply_activation_tensor(x, &layer.activation)?; } } Ok(x) } /// Get reference to physics loss function pub fn physics_loss(&self) -> &dyn PhysicsLoss { self.physics_loss.as_ref() } } impl std::fmt::Debug for PINN { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("PINN") .field("network", &self.network) .field("device", &self.device) .field("physics_loss", &"") .field( "conservation_losses", &format!("{} losses", self.conservation_losses.len()), ) .field("config", &self.config) .finish() } } impl PINNBuilder { /// Create a new PINN builder #[must_use] pub fn new() -> Self { Self { device: None, layers: Vec::new(), activation: ActivationType::Tanh, physics_loss: None, conservation_losses: Vec::new(), config: TrainingConfig::default(), } } /// Set the computation device #[must_use] pub fn device(mut self, device: &Device) -> Self { self.device = Some(device.clone()); self } /// Set the network layer configuration #[must_use] pub fn layers(mut self, layers: Vec) -> Self { self.layers = layers; self } /// Set the activation function #[must_use] pub fn activation(mut self, activation: ActivationType) -> Self { self.activation = activation; self } /// Set the physics loss function #[must_use] pub fn physics_loss(mut self, loss: Box) -> Self { self.physics_loss = Some(loss); self } /// Add a conservation law constraint #[must_use] pub fn conservation_loss(mut self, loss: Box) -> Self { self.conservation_losses.push(loss); self } /// Set the training configuration #[must_use] pub fn config(mut self, config: TrainingConfig) -> Self { self.config = config; self } /// Build the PINN pub fn build(self) -> Result { PINN::from_builder(self) } } impl Default for PINNBuilder { fn default() -> Self { Self::new() } } impl LinearLayer { /// Create a new linear layer fn new( input_size: usize, output_size: usize, activation: &ActivationType, device: &Device, ) -> Result { // Xavier/Glorot initialization let std_dev = (2.0 / (input_size + output_size) as f64).sqrt(); let weights = Variable::new( Tensor::randn(&[input_size, output_size], device)?.mul_scalar(std_dev as f32)?, true, // requires_grad ); let bias = Variable::new(Tensor::zeros([output_size], device)?, true); Ok(Self { weights, bias, activation: activation.clone(), }) } /// Forward pass through the layer fn forward(&self, input: Variable) -> Result { let output = input .matmul(&self.weights) .map_err(|e| ScienceError::computation(e.to_string()))?; output .add(&self.bias) .map_err(|e| ScienceError::computation(e.to_string())) } } /// Apply activation function to a variable fn apply_activation(input: Variable, activation: &ActivationType) -> Result { use crate::variable_extensions::VariableExt; match activation { ActivationType::Tanh => input.tanh(), ActivationType::Sigmoid => input.sigmoid(), ActivationType::Sin => input.sin(), ActivationType::Swish => { let sigmoid = input.sigmoid()?; input.multiply(&sigmoid).map_err(std::convert::Into::into) } ActivationType::GELU => { // GELU approximation: 0.5 * x * (1 + tanh(sqrt(2/π) * (x + 0.044715 * x^3))) let x_cubed = input.multiply(&input)?.multiply(&input)?; let inner = input.add(&x_cubed.multiply_scalar(0.044715)?)?; let inner = inner.multiply_scalar((2.0 / std::f64::consts::PI).sqrt() as f32)?; let tanh_part = inner.tanh()?; let one_plus_tanh = tanh_part.add_scalar(1.0)?; let result = input.multiply(&one_plus_tanh)?; result .multiply_scalar(0.5) .map_err(std::convert::Into::into) } ActivationType::ReLU => VariableExt::relu(&input), ActivationType::Learnable => { // Simple learnable activation: a * tanh(b * x) // For now, use standard tanh (parameters would need to be learned) VariableExt::tanh(&input) } } } /// Apply activation function to a tensor (non-gradient version) fn apply_activation_tensor(input: Tensor, activation: &ActivationType) -> Result { match activation { ActivationType::Tanh => input.tanh().map_err(std::convert::Into::into), ActivationType::Sigmoid => input.sigmoid().map_err(std::convert::Into::into), ActivationType::Sin => input.sin().map_err(std::convert::Into::into), ActivationType::Swish => { let sigmoid = input .sigmoid() .map_err(|e: TensorError| -> ScienceError { e.into() })?; input.mul(&sigmoid).map_err(std::convert::Into::into) } ActivationType::GELU => { // GELU approximation: 0.5 * x * (1 + tanh(sqrt(2/π) * (x + 0.044715 * x^3))) let x_squared = input .mul(&input) .map_err(|e: TensorError| -> ScienceError { e.into() })?; let x_cubed = x_squared .mul(&input) .map_err(|e: TensorError| -> ScienceError { e.into() })?; let inner = input .add( &x_cubed .mul_scalar(0.044715) .map_err(|e: TensorError| -> ScienceError { e.into() })?, ) .map_err(|e: TensorError| -> ScienceError { e.into() })?; let inner = inner .mul_scalar((2.0 / std::f64::consts::PI).sqrt() as f32) .map_err(|e: TensorError| -> ScienceError { e.into() })?; let tanh_part = inner .tanh() .map_err(|e: TensorError| -> ScienceError { e.into() })?; let one = Tensor::ones(input.shape().dims(), input.device()) .map_err(|e: TensorError| -> ScienceError { e.into() })?; let one_plus_tanh = one .add(&tanh_part) .map_err(|e: TensorError| -> ScienceError { e.into() })?; let result = input .mul(&one_plus_tanh) .map_err(|e: TensorError| -> ScienceError { e.into() })?; result.mul_scalar(0.5).map_err(std::convert::Into::into) } ActivationType::ReLU => input.relu().map_err(std::convert::Into::into), ActivationType::Learnable => { // Simple learnable activation: for now, use standard tanh input.tanh().map_err(std::convert::Into::into) } } } /// Create input tensor from coordinate pairs fn create_input_tensor(inputs: &[(f64, f64)], device: &Device) -> Result { let data: Vec = inputs .iter() .flat_map(|(x, t)| vec![*x as f32, *t as f32]) .collect(); Ok(Tensor::from_slice(&data, &[inputs.len(), 2], device)?) } #[cfg(test)] mod tests { use super::*; use crate::physics::HeatEquation; #[tokio::test] async fn test_pinn_creation() -> Result<()> { let device = Device::cpu(); let heat_eq = HeatEquation::new(0.1); let pinn = PINN::builder() .device(&device) .layers(vec![2, 64, 64, 1]) .physics_loss(Box::new(heat_eq)) .build()?; assert_eq!(pinn.network.len(), 3); Ok(()) } #[tokio::test] async fn test_pinn_forward_pass() -> Result<()> { let device = Device::cpu(); let heat_eq = HeatEquation::new(0.1); let pinn = PINN::builder() .device(&device) .layers(vec![2, 32, 1]) .physics_loss(Box::new(heat_eq)) .build()?; // Create input tensor with shape [2, 2] for 2 samples with 2 features each let inputs = Tensor::from_slice(&[0.5_f32, 1.0, 0.3, 0.8], &[2, 2], &device)?; let outputs = pinn.predict(&inputs).await?; // Output shape should be [2, 1] for 2 samples with 1 output each assert_eq!(outputs.shape().dims(), &[2, 1]); Ok(()) } }