//! Conservation Laws for Physics-Informed Neural Networks //! //! This module implements various conservation laws that can be enforced //! in PINNs to ensure physical consistency of the learned solutions. use crate::Tensor; use crate::error::{ConservationLaw, Result, ScienceError}; use crate::variable_extensions::VariableExt; use async_trait::async_trait; use rtx_autograd::Variable; use serde::{Deserialize, Serialize}; use std::collections::HashMap; /// Trait for conservation law constraints in PINNs #[async_trait] pub trait ConservationLoss: Send + Sync { /// Compute the conservation law violation async fn compute_loss( &self, coordinates: &Tensor, // Input coordinates [x, t, ...] solution: &Variable, // Network output u(x,t) du_dx: &Variable, // ∂u/∂x du_dt: &Variable, // ∂u/∂t ) -> Result; /// Get the name of the conservation law fn name(&self) -> &str; /// Get the conservation law type fn law_type(&self) -> ConservationLaw; /// Get tolerance for violation detection fn tolerance(&self) -> f64; /// Validate conservation at given points async fn validate_conservation( &self, coordinates: &Tensor, solution: &Variable, du_dx: &Variable, du_dt: &Variable, ) -> Result; } /// Conservation law validation result #[derive(Debug, Clone)] pub struct ValidationResult { /// Maximum violation magnitude pub max_violation: f64, /// Mean violation magnitude pub mean_violation: f64, /// Points where tolerance is exceeded pub violation_points: Vec, /// Whether conservation is satisfied within tolerance pub is_satisfied: bool, /// Detailed diagnostics pub diagnostics: HashMap, } /// Mass conservation law: ∂ρ/∂t + ∇·(ρv) = 0 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct MassConservation { /// Density field identifier pub density_field: String, /// Velocity field identifiers [u, v, w] pub velocity_fields: Vec, /// Tolerance for conservation violation pub tolerance: f64, /// Weight in loss function pub weight: f64, } /// Momentum conservation law: ∂(ρv)/∂t + ∇·(ρv⊗v) = -∇p + μ∇²v + f #[derive(Debug, Clone, Serialize, Deserialize)] pub struct MomentumConservation { /// Density field pub density_field: String, /// Velocity components pub velocity_fields: Vec, /// Pressure field pub pressure_field: String, /// Viscosity coefficient pub viscosity: f64, /// External force fields pub force_fields: Vec, /// Tolerance for conservation violation pub tolerance: f64, /// Weight in loss function pub weight: f64, } /// Energy conservation law: ∂E/∂t + ∇·(Ev + pv - k∇T) = Q #[derive(Debug, Clone, Serialize, Deserialize)] pub struct EnergyConservation { /// Total energy field pub energy_field: String, /// Velocity fields pub velocity_fields: Vec, /// Pressure field pub pressure_field: String, /// Temperature field pub temperature_field: String, /// Thermal conductivity pub thermal_conductivity: f64, /// Heat source term pub heat_source: f64, /// Tolerance for conservation violation pub tolerance: f64, /// Weight in loss function pub weight: f64, } /// Angular momentum conservation: L = r × p = constant #[derive(Debug, Clone, Serialize, Deserialize)] pub struct AngularMomentumConservation { /// Position fields [x, y, z] pub position_fields: Vec, /// Momentum fields [px, py, pz] pub momentum_fields: Vec, /// Reference point for angular momentum calculation pub reference_point: Vec, /// Tolerance for conservation violation pub tolerance: f64, /// Weight in loss function pub weight: f64, } /// Charge conservation: ∂ρ/∂t + ∇·J = 0 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ChargeConservation { /// Charge density field pub charge_density_field: String, /// Current density fields [Jx, Jy, Jz] pub current_density_fields: Vec, /// Tolerance for conservation violation pub tolerance: f64, /// Weight in loss function pub weight: f64, } /// Probability conservation (quantum mechanics): ∫|ψ|²dx = 1 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ProbabilityConservation { /// Wavefunction field (complex) pub wavefunction_field: String, /// Integration domain bounds pub domain_bounds: Vec<(f64, f64)>, /// Tolerance for conservation violation pub tolerance: f64, /// Weight in loss function pub weight: f64, } /// Conservation validator for comprehensive checking pub struct ConservationValidator { /// List of conservation laws to check conservation_laws: Vec>, /// Global tolerance settings global_tolerance: f64, /// Validation frequency (every N training steps) validation_frequency: usize, /// Current validation step current_step: usize, } impl std::fmt::Debug for ConservationValidator { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("ConservationValidator") .field("num_conservation_laws", &self.conservation_laws.len()) .field("global_tolerance", &self.global_tolerance) .field("validation_frequency", &self.validation_frequency) .field("current_step", &self.current_step) .finish() } } impl MassConservation { /// Create a new mass conservation law #[must_use] pub fn new(tolerance: f64) -> Self { Self { density_field: "rho".to_string(), velocity_fields: vec!["u".to_string(), "v".to_string()], tolerance, weight: 1.0, } } /// Set velocity field names #[must_use] pub fn with_velocity_fields(mut self, fields: Vec) -> Self { self.velocity_fields = fields; self } /// Set loss weight #[must_use] pub fn with_weight(mut self, weight: f64) -> Self { self.weight = weight; self } } #[async_trait] impl ConservationLoss for MassConservation { async fn compute_loss( &self, _coordinates: &Tensor, _solution: &Variable, du_dx: &Variable, du_dt: &Variable, ) -> Result { // For simplified mass conservation: ∂ρ/∂t + ∂(ρu)/∂x = 0 // Assuming solution represents density ρ and we have velocity field // ∂ρ/∂t term let drho_dt = du_dt.clone(); // ∂(ρu)/∂x term - simplified to ρ∂u/∂x + u∂ρ/∂x // For this example, assume constant velocity or treat solution as ρu let drhou_dx = du_dx.clone(); // Mass conservation residual: ∂ρ/∂t + ∂(ρu)/∂x = 0 let residual = drho_dt.add(&drhou_dx)?; // Return mean squared residual as f32 let squared_residual = residual.multiply(&residual)?; let loss_var = squared_residual.mean_square()?; let loss_value = loss_var.value().to_scalar::()?; Ok(loss_value) } fn name(&self) -> &'static str { "Mass Conservation" } fn law_type(&self) -> ConservationLaw { ConservationLaw::Mass } fn tolerance(&self) -> f64 { self.tolerance } async fn validate_conservation( &self, coordinates: &Tensor, solution: &Variable, du_dx: &Variable, du_dt: &Variable, ) -> Result { // Compute conservation residual let residual = self .compute_loss(coordinates, solution, du_dx, du_dt) .await?; let residual_values = [residual]; // f32 scalar converted to vec let violations: Vec = residual_values .iter() .map(|&x| f64::from(x.abs())) .collect(); let max_violation = violations.iter().copied().fold(0.0, f64::max); let mean_violation = violations.iter().sum::() / violations.len() as f64; let violation_points: Vec = violations .iter() .enumerate() .filter_map(|(i, &v)| if v > self.tolerance { Some(i) } else { None }) .collect(); let is_satisfied = max_violation <= self.tolerance; let mut diagnostics = HashMap::new(); diagnostics.insert("max_violation".to_string(), max_violation); diagnostics.insert("mean_violation".to_string(), mean_violation); diagnostics.insert( "violation_rate".to_string(), violation_points.len() as f64 / violations.len() as f64, ); Ok(ValidationResult { max_violation, mean_violation, violation_points, is_satisfied, diagnostics, }) } } impl MomentumConservation { /// Create a new momentum conservation law #[must_use] pub fn new(viscosity: f64, tolerance: f64) -> Self { Self { density_field: "rho".to_string(), velocity_fields: vec!["u".to_string(), "v".to_string()], pressure_field: "p".to_string(), viscosity, force_fields: Vec::new(), tolerance, weight: 1.0, } } /// Add external force field #[must_use] pub fn with_force_field(mut self, force_field: String) -> Self { self.force_fields.push(force_field); self } } #[async_trait] impl ConservationLoss for MomentumConservation { async fn compute_loss( &self, _coordinates: &Tensor, _solution: &Variable, du_dx: &Variable, du_dt: &Variable, ) -> Result { // Simplified momentum conservation: ∂(ρu)/∂t + ∂p/∂x = μ∂²u/∂x² // Assuming solution represents momentum ρu let momentum_time_derivative = du_dt.clone(); // For now, assume pressure gradient is represented by some component // In a full implementation, this would be computed from pressure field let pressure_gradient = du_dx.multiply_scalar(0.1)?; // Placeholder // Viscous term (simplified) let viscous_term = du_dx.multiply_scalar(self.viscosity as f32)?; // Momentum conservation residual let residual = momentum_time_derivative .add(&pressure_gradient)? .subtract(&viscous_term)?; let squared_residual = residual.multiply(&residual)?; let loss_var = squared_residual.mean_square()?; let loss_value = loss_var.value().to_scalar::()?; Ok(loss_value) } fn name(&self) -> &'static str { "Momentum Conservation" } fn law_type(&self) -> ConservationLaw { ConservationLaw::Momentum } fn tolerance(&self) -> f64 { self.tolerance } async fn validate_conservation( &self, coordinates: &Tensor, solution: &Variable, du_dx: &Variable, du_dt: &Variable, ) -> Result { let residual = self .compute_loss(coordinates, solution, du_dx, du_dt) .await?; let residual_values = [residual]; // f32 scalar converted to vec let violations: Vec = residual_values .iter() .map(|&x| f64::from(x.abs())) .collect(); let max_violation = violations.iter().copied().fold(0.0, f64::max); let mean_violation = violations.iter().sum::() / violations.len() as f64; let violation_points: Vec = violations .iter() .enumerate() .filter_map(|(i, &v)| if v > self.tolerance { Some(i) } else { None }) .collect(); let is_satisfied = max_violation <= self.tolerance; let mut diagnostics = HashMap::new(); diagnostics.insert("max_violation".to_string(), max_violation); diagnostics.insert("mean_violation".to_string(), mean_violation); diagnostics.insert("viscosity".to_string(), self.viscosity); Ok(ValidationResult { max_violation, mean_violation, violation_points, is_satisfied, diagnostics, }) } } impl EnergyConservation { /// Create a new energy conservation law #[must_use] pub fn new(thermal_conductivity: f64, tolerance: f64) -> Self { Self { energy_field: "E".to_string(), velocity_fields: vec!["u".to_string(), "v".to_string()], pressure_field: "p".to_string(), temperature_field: "T".to_string(), thermal_conductivity, heat_source: 0.0, tolerance, weight: 1.0, } } /// Set heat source term #[must_use] pub fn with_heat_source(mut self, heat_source: f64) -> Self { self.heat_source = heat_source; self } } #[async_trait] impl ConservationLoss for EnergyConservation { async fn compute_loss( &self, _coordinates: &Tensor, _solution: &Variable, du_dx: &Variable, du_dt: &Variable, ) -> Result { // Simplified energy conservation: ∂E/∂t = -∂q/∂x + Q // where q is heat flux and Q is heat source let energy_time_derivative = du_dt.clone(); // Heat flux term (simplified as thermal conductivity * temperature gradient) let heat_flux_gradient = du_dx.multiply_scalar(self.thermal_conductivity as f32)?; // Heat source term let du_dt_tensor = du_dt.value(); let source_shape = &[du_dt_tensor.shape().dims()[0]]; let source_tensor = Tensor::full(source_shape, self.heat_source as f32, du_dt_tensor.device())?; let source_term = Variable::constant(source_tensor); // Energy conservation residual let residual = energy_time_derivative .add(&heat_flux_gradient)? .subtract(&source_term)?; let squared_residual = residual.multiply(&residual)?; let loss_var = squared_residual.mean_square()?; let loss_value = loss_var.value().to_scalar::()?; Ok(loss_value) } fn name(&self) -> &'static str { "Energy Conservation" } fn law_type(&self) -> ConservationLaw { ConservationLaw::Energy } fn tolerance(&self) -> f64 { self.tolerance } async fn validate_conservation( &self, coordinates: &Tensor, solution: &Variable, du_dx: &Variable, du_dt: &Variable, ) -> Result { let residual = self .compute_loss(coordinates, solution, du_dx, du_dt) .await?; let residual_values = [residual]; // f32 scalar converted to vec let violations: Vec = residual_values .iter() .map(|&x| f64::from(x.abs())) .collect(); let max_violation = violations.iter().copied().fold(0.0, f64::max); let mean_violation = violations.iter().sum::() / violations.len() as f64; let violation_points: Vec = violations .iter() .enumerate() .filter_map(|(i, &v)| if v > self.tolerance { Some(i) } else { None }) .collect(); let is_satisfied = max_violation <= self.tolerance; let mut diagnostics = HashMap::new(); diagnostics.insert("max_violation".to_string(), max_violation); diagnostics.insert("mean_violation".to_string(), mean_violation); diagnostics.insert( "thermal_conductivity".to_string(), self.thermal_conductivity, ); diagnostics.insert("heat_source".to_string(), self.heat_source); Ok(ValidationResult { max_violation, mean_violation, violation_points, is_satisfied, diagnostics, }) } } impl ConservationValidator { /// Create a new conservation validator #[must_use] pub fn new(global_tolerance: f64, validation_frequency: usize) -> Self { Self { conservation_laws: Vec::new(), global_tolerance, validation_frequency, current_step: 0, } } /// Add a conservation law to validate #[must_use] pub fn add_conservation_law(mut self, law: Box) -> Self { self.conservation_laws.push(law); self } /// Validate all conservation laws pub async fn validate_all( &mut self, coordinates: &Tensor, solution: &Variable, du_dx: &Variable, du_dt: &Variable, ) -> Result> { self.current_step += 1; // Only validate at specified frequency if !self.current_step.is_multiple_of(self.validation_frequency) { return Ok(Vec::new()); } let mut results = Vec::new(); for law in &self.conservation_laws { let result = law .validate_conservation(coordinates, solution, du_dx, du_dt) .await?; // Log critical violations if !result.is_satisfied && result.max_violation > self.global_tolerance * 10.0 { tracing::warn!( "Critical conservation violation in {}: max = {:.2e}, tolerance = {:.2e}", law.name(), result.max_violation, law.tolerance() ); } results.push(result); } Ok(results) } /// Get summary of all conservation laws #[must_use] pub fn summary(&self) -> HashMap { let mut summary = HashMap::new(); summary.insert( "num_laws".to_string(), self.conservation_laws.len().to_string(), ); summary.insert( "global_tolerance".to_string(), self.global_tolerance.to_string(), ); summary.insert( "validation_frequency".to_string(), self.validation_frequency.to_string(), ); summary.insert("current_step".to_string(), self.current_step.to_string()); for (i, law) in self.conservation_laws.iter().enumerate() { summary.insert(format!("law_{i}_name"), law.name().to_string()); summary.insert(format!("law_{i}_tolerance"), law.tolerance().to_string()); } summary } } /// Factory function for creating conservation laws pub fn create_conservation_law( law_type: ConservationLaw, tolerance: f64, parameters: &HashMap, ) -> Result> { match law_type { ConservationLaw::Mass => Ok(Box::new(MassConservation::new(tolerance))), ConservationLaw::Momentum => { let viscosity = parameters.get("viscosity").copied().unwrap_or(1e-3); Ok(Box::new(MomentumConservation::new(viscosity, tolerance))) } ConservationLaw::Energy => { let thermal_conductivity = parameters .get("thermal_conductivity") .copied() .unwrap_or(0.1); let mut energy_law = EnergyConservation::new(thermal_conductivity, tolerance); if let Some(&heat_source) = parameters.get("heat_source") { energy_law = energy_law.with_heat_source(heat_source); } Ok(Box::new(energy_law)) } _ => Err(ScienceError::physics( format!("Unsupported conservation law: {law_type:?}"), crate::error::PhysicsDomain::FluidDynamics, )), } } #[cfg(test)] mod tests { use super::*; use crate::Device; #[tokio::test] async fn test_mass_conservation() -> Result<()> { let device = Device::cpu(); let conservation = MassConservation::new(1e-6); // Create dummy data let coords = Tensor::zeros(&[10, 2], &device)?; let solution = Variable::new(Tensor::ones(&[10], &device)?, true); let du_dx = Variable::new(Tensor::zeros(&[10], &device)?, false); let du_dt = Variable::new(Tensor::zeros(&[10], &device)?, false); let loss = conservation .compute_loss(&coords, &solution, &du_dx, &du_dt) .await?; assert!(loss >= 0.0); Ok(()) } #[test] fn test_conservation_validator() { let validator = ConservationValidator::new(1e-6, 10); assert_eq!(validator.conservation_laws.len(), 0); let summary = validator.summary(); assert_eq!(summary["num_laws"], "0"); } #[test] fn test_conservation_law_factory() -> Result<()> { let mut params = HashMap::new(); params.insert("viscosity".to_string(), 1e-3); let law = create_conservation_law(ConservationLaw::Momentum, 1e-6, ¶ms)?; assert_eq!(law.name(), "Momentum Conservation"); Ok(()) } }