Initial commit
This commit is contained in:
@@ -0,0 +1,531 @@
|
||||
//! Training utilities for Physics-Informed Neural Networks
|
||||
|
||||
use crate::error::{Result, ScienceError};
|
||||
use crate::physics::{BoundaryConditions, LossComponents, PINN};
|
||||
use rtx_autograd::Variable;
|
||||
use rtx_tensor::{Device, Tensor};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::VecDeque;
|
||||
|
||||
/// Training configuration for PINNs
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TrainingConfig {
|
||||
/// Learning rate
|
||||
pub learning_rate: f64,
|
||||
/// Physics loss weight
|
||||
pub physics_weight: f64,
|
||||
/// Boundary condition loss weight
|
||||
pub boundary_weight: f64,
|
||||
/// Conservation law loss weights
|
||||
pub conservation_weights: Vec<f64>,
|
||||
/// Convergence tolerance
|
||||
pub convergence_tolerance: f64,
|
||||
/// Early stopping patience
|
||||
pub patience: usize,
|
||||
}
|
||||
|
||||
impl Default for TrainingConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
learning_rate: 1e-3,
|
||||
physics_weight: 1.0,
|
||||
boundary_weight: 100.0,
|
||||
conservation_weights: vec![1.0],
|
||||
convergence_tolerance: 1e-6,
|
||||
patience: 100,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// PINN trainer
|
||||
pub struct PINNTrainer {
|
||||
config: TrainingConfig,
|
||||
}
|
||||
|
||||
impl PINNTrainer {
|
||||
#[must_use]
|
||||
pub fn new(config: TrainingConfig) -> Self {
|
||||
Self { config }
|
||||
}
|
||||
|
||||
pub async fn train(
|
||||
&self,
|
||||
pinn: &mut PINN,
|
||||
boundary_data: BoundaryConditions,
|
||||
epochs: usize,
|
||||
) -> Result<()> {
|
||||
use std::collections::VecDeque;
|
||||
|
||||
tracing::info!("Starting PINN training for {} epochs", epochs);
|
||||
|
||||
// Initialize training state
|
||||
let mut optimizer = AdamOptimizer::new(self.config.learning_rate);
|
||||
let mut loss_history = VecDeque::with_capacity(self.config.patience);
|
||||
let mut best_loss = f64::INFINITY;
|
||||
let mut patience_counter = 0;
|
||||
let lr_scheduler = CosineAnnealingScheduler::new(
|
||||
self.config.learning_rate,
|
||||
epochs,
|
||||
0.01 * self.config.learning_rate,
|
||||
);
|
||||
|
||||
// Generate training points using Latin Hypercube Sampling
|
||||
let training_points = self.generate_training_points(&boundary_data, 1000)?;
|
||||
let boundary_points = self.sample_boundary_points(&boundary_data, 200)?;
|
||||
|
||||
// Convert points to tensors
|
||||
let training_tensor = self.points_to_tensor(&training_points)?;
|
||||
let _boundary_tensor = self.points_to_tensor(&boundary_points)?;
|
||||
|
||||
for epoch in 0..epochs {
|
||||
let epoch_start = std::time::Instant::now();
|
||||
|
||||
// Forward pass - compute network output and gradients
|
||||
let (outputs, du_dx, du_dt) = pinn.forward_with_gradients(&training_tensor).await?;
|
||||
|
||||
// Create gradients struct
|
||||
// For now, we'll compute second derivatives when needed
|
||||
let gradients = PhysicsGradients {
|
||||
du_dx: du_dx.clone(),
|
||||
du_dt: du_dt.clone(),
|
||||
d2u_dx2: du_dx.clone(), // Placeholder - should compute actual second derivative
|
||||
d2u_dt2: du_dt.clone(), // Placeholder - should compute actual second derivative
|
||||
d2u_dxdt: du_dx.clone(), // Placeholder - should compute actual second derivative
|
||||
};
|
||||
|
||||
// Compute loss components
|
||||
let mut total_loss = 0.0;
|
||||
let mut loss_components = LossComponents {
|
||||
physics: 0.0,
|
||||
boundary: 0.0,
|
||||
initial: 0.0,
|
||||
conservation: Vec::new(),
|
||||
data: 0.0,
|
||||
total: 0.0,
|
||||
iteration: epoch,
|
||||
};
|
||||
|
||||
// 1. Physics loss (PDE residual)
|
||||
let physics_residual = self
|
||||
.compute_physics_residual(pinn, &training_points, &outputs, &gradients)
|
||||
.await?;
|
||||
// Compute mean squared physics residual
|
||||
let physics_squared = physics_residual.multiply(&physics_residual)?;
|
||||
let physics_loss_var = physics_squared.mean_square()?;
|
||||
loss_components.physics = f64::from(physics_loss_var.value().to_scalar::<f32>()?);
|
||||
total_loss += self.config.physics_weight * loss_components.physics;
|
||||
|
||||
// 2. Boundary condition loss
|
||||
let boundary_loss = self
|
||||
.compute_boundary_loss(pinn, &boundary_points, &boundary_data)
|
||||
.await?;
|
||||
loss_components.boundary = boundary_loss as f64;
|
||||
total_loss += self.config.boundary_weight * boundary_loss as f64;
|
||||
|
||||
// 3. Initial condition loss (if time-dependent)
|
||||
if self.is_time_dependent(&boundary_data) {
|
||||
let initial_loss = self
|
||||
.compute_initial_condition_loss(pinn, &boundary_data)
|
||||
.await?;
|
||||
loss_components.initial = initial_loss as f64;
|
||||
total_loss += self.config.boundary_weight * initial_loss as f64;
|
||||
}
|
||||
|
||||
// 4. Conservation law losses
|
||||
// For now skip conservation losses as they would need to be accessed differently
|
||||
// This would need a refactor to expose conservation_losses from PINN
|
||||
|
||||
loss_components.total = total_loss;
|
||||
|
||||
// Backward pass and optimization
|
||||
let loss_tensor = self.create_loss_tensor(total_loss)?;
|
||||
let _gradients = loss_tensor.backward(); // Returns HashMap, not Result
|
||||
|
||||
// Gradient clipping for stability
|
||||
pinn.clip_gradients(1.0);
|
||||
|
||||
// Optimizer step
|
||||
optimizer.step(pinn.parameters()).await?;
|
||||
pinn.zero_gradients();
|
||||
|
||||
// Update learning rate
|
||||
let new_lr = lr_scheduler.step(epoch);
|
||||
optimizer.set_learning_rate(new_lr);
|
||||
|
||||
// Track convergence
|
||||
loss_history.push_back(total_loss);
|
||||
if loss_history.len() > self.config.patience {
|
||||
loss_history.pop_front();
|
||||
}
|
||||
|
||||
// Early stopping check
|
||||
if total_loss < best_loss {
|
||||
best_loss = total_loss;
|
||||
patience_counter = 0;
|
||||
// Save best model state
|
||||
pinn.save_state().await?;
|
||||
} else {
|
||||
patience_counter += 1;
|
||||
}
|
||||
|
||||
// Check convergence
|
||||
let converged = self.check_convergence(&loss_history, loss_components.total)?;
|
||||
|
||||
// Adaptive weight scheduling
|
||||
if epoch > 0 && epoch % 100 == 0 {
|
||||
self.update_loss_weights(pinn, &loss_components).await?;
|
||||
}
|
||||
|
||||
// Logging
|
||||
if epoch % 10 == 0 || converged {
|
||||
let epoch_time = epoch_start.elapsed();
|
||||
tracing::info!(
|
||||
"Epoch {}/{}: Loss={:.6e} (Physics={:.6e}, Boundary={:.6e}, Initial={:.6e}) LR={:.6e} Time={:.2?}",
|
||||
epoch + 1,
|
||||
epochs,
|
||||
loss_components.total,
|
||||
loss_components.physics,
|
||||
loss_components.boundary,
|
||||
loss_components.initial,
|
||||
new_lr,
|
||||
epoch_time
|
||||
);
|
||||
}
|
||||
|
||||
// Early stopping
|
||||
if patience_counter >= self.config.patience && epoch > 100 {
|
||||
tracing::info!(
|
||||
"Early stopping triggered at epoch {} (patience={})",
|
||||
epoch + 1,
|
||||
patience_counter
|
||||
);
|
||||
break;
|
||||
}
|
||||
|
||||
if converged {
|
||||
tracing::info!("Convergence achieved at epoch {}", epoch + 1);
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
// Restore best model if early stopping occurred
|
||||
if patience_counter > 0 {
|
||||
let checkpoint = pinn.save_state().await?;
|
||||
pinn.restore_state(checkpoint).await?;
|
||||
}
|
||||
|
||||
tracing::info!("PINN training completed. Final loss: {:.6e}", best_loss);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Generate training points using Latin Hypercube Sampling for better space coverage
|
||||
fn generate_training_points(
|
||||
&self,
|
||||
boundary_data: &BoundaryConditions,
|
||||
num_points: usize,
|
||||
) -> Result<Vec<(f64, f64)>> {
|
||||
use rand::Rng;
|
||||
let mut rng = rand::thread_rng();
|
||||
let mut points = Vec::with_capacity(num_points);
|
||||
|
||||
// Get domain bounds from boundary conditions
|
||||
let (x_min, x_max) = boundary_data.spatial_bounds();
|
||||
let (t_min, t_max) = boundary_data.temporal_bounds();
|
||||
|
||||
// Latin Hypercube Sampling
|
||||
let mut x_samples: Vec<f64> = (0..num_points)
|
||||
.map(|i| (i as f64 + rng.r#gen::<f64>()) / num_points as f64)
|
||||
.collect();
|
||||
let mut t_samples: Vec<f64> = (0..num_points)
|
||||
.map(|i| (i as f64 + rng.r#gen::<f64>()) / num_points as f64)
|
||||
.collect();
|
||||
|
||||
// Shuffle to break correlation
|
||||
use rand::seq::SliceRandom;
|
||||
x_samples.shuffle(&mut rng);
|
||||
t_samples.shuffle(&mut rng);
|
||||
|
||||
// Scale to actual domain
|
||||
for i in 0..num_points {
|
||||
let x = x_min + x_samples[i] * (x_max - x_min);
|
||||
let t = t_min + t_samples[i] * (t_max - t_min);
|
||||
points.push((x, t));
|
||||
}
|
||||
|
||||
Ok(points)
|
||||
}
|
||||
|
||||
/// Sample boundary points for boundary condition enforcement
|
||||
fn sample_boundary_points(
|
||||
&self,
|
||||
boundary_data: &BoundaryConditions,
|
||||
num_points: usize,
|
||||
) -> Result<Vec<(f64, f64)>> {
|
||||
let mut points = Vec::with_capacity(num_points);
|
||||
let (x_min, x_max) = boundary_data.spatial_bounds();
|
||||
let (t_min, t_max) = boundary_data.temporal_bounds();
|
||||
|
||||
// Sample points on all boundaries
|
||||
let points_per_boundary = num_points / 4;
|
||||
|
||||
// Left boundary (x = x_min)
|
||||
for i in 0..points_per_boundary {
|
||||
let t = t_min + (i as f64 / points_per_boundary as f64) * (t_max - t_min);
|
||||
points.push((x_min, t));
|
||||
}
|
||||
|
||||
// Right boundary (x = x_max)
|
||||
for i in 0..points_per_boundary {
|
||||
let t = t_min + (i as f64 / points_per_boundary as f64) * (t_max - t_min);
|
||||
points.push((x_max, t));
|
||||
}
|
||||
|
||||
// Bottom boundary (t = t_min)
|
||||
for i in 0..points_per_boundary {
|
||||
let x = x_min + (i as f64 / points_per_boundary as f64) * (x_max - x_min);
|
||||
points.push((x, t_min));
|
||||
}
|
||||
|
||||
// Top boundary (t = t_max)
|
||||
for i in 0..points_per_boundary {
|
||||
let x = x_min + (i as f64 / points_per_boundary as f64) * (x_max - x_min);
|
||||
points.push((x, t_max));
|
||||
}
|
||||
|
||||
Ok(points)
|
||||
}
|
||||
|
||||
/// Compute physics residual (PDE satisfaction)
|
||||
async fn compute_physics_residual(
|
||||
&self,
|
||||
pinn: &PINN,
|
||||
training_points: &[(f64, f64)],
|
||||
outputs: &Variable,
|
||||
gradients: &PhysicsGradients,
|
||||
) -> Result<Variable> {
|
||||
// This would call the physics loss function to compute PDE residual
|
||||
pinn.physics_loss()
|
||||
.compute_residual(
|
||||
&self.points_to_tensor(training_points)?,
|
||||
outputs,
|
||||
&gradients.du_dx,
|
||||
&gradients.du_dt,
|
||||
&gradients.d2u_dx2,
|
||||
&gradients.d2u_dt2,
|
||||
&gradients.d2u_dxdt,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Compute boundary condition loss
|
||||
async fn compute_boundary_loss(
|
||||
&self,
|
||||
pinn: &PINN,
|
||||
boundary_points: &[(f64, f64)],
|
||||
boundary_data: &BoundaryConditions,
|
||||
) -> Result<f64> {
|
||||
let mut total_loss = 0.0;
|
||||
let mut point_count = 0;
|
||||
|
||||
for (x, t) in boundary_points {
|
||||
let _input = [*x, *t];
|
||||
let input_tensor =
|
||||
Tensor::from_slice(&[*x as f32, *t as f32], &[1, 2], &Device::Cuda(0))?;
|
||||
let predicted_tensor = pinn.predict(&input_tensor).await?;
|
||||
let predicted = predicted_tensor.to_cpu()?;
|
||||
|
||||
let expected_value = boundary_data.get_value_at(*x, *t);
|
||||
let error = (f64::from(predicted[0]) - expected_value).powi(2);
|
||||
total_loss += error;
|
||||
point_count += 1;
|
||||
}
|
||||
|
||||
Ok(if point_count > 0 {
|
||||
total_loss / f64::from(point_count)
|
||||
} else {
|
||||
0.0
|
||||
})
|
||||
}
|
||||
|
||||
/// Compute initial condition loss
|
||||
async fn compute_initial_condition_loss(
|
||||
&self,
|
||||
pinn: &PINN,
|
||||
boundary_data: &BoundaryConditions,
|
||||
) -> Result<f64> {
|
||||
let (x_min, x_max) = boundary_data.spatial_bounds();
|
||||
let t_initial = boundary_data.temporal_bounds().0;
|
||||
|
||||
let mut total_loss = 0.0;
|
||||
let num_points = 100;
|
||||
|
||||
for i in 0..num_points {
|
||||
let x = x_min + (f64::from(i) / f64::from(num_points)) * (x_max - x_min);
|
||||
let input_tensor =
|
||||
Tensor::from_slice(&[x as f32, t_initial as f32], &[1, 2], &Device::Cuda(0))?;
|
||||
let predicted_tensor = pinn.predict(&input_tensor).await?;
|
||||
let predicted_vec = predicted_tensor.to_cpu()?;
|
||||
|
||||
let initial_value = boundary_data.get_initial_value_at(x);
|
||||
let error = (f64::from(predicted_vec[0]) - initial_value).powi(2);
|
||||
total_loss += error;
|
||||
}
|
||||
|
||||
Ok(total_loss / f64::from(num_points))
|
||||
}
|
||||
|
||||
/// Check if the problem is time-dependent
|
||||
fn is_time_dependent(&self, boundary_data: &BoundaryConditions) -> bool {
|
||||
let (t_min, t_max) = boundary_data.temporal_bounds();
|
||||
t_max > t_min
|
||||
}
|
||||
|
||||
/// Create loss tensor from scalar value
|
||||
fn create_loss_tensor(&self, loss_value: f64) -> Result<Variable> {
|
||||
let tensor = Tensor::scalar(loss_value as f32, rtx_tensor::DType::F32, &Device::Cuda(0))?;
|
||||
Ok(Variable::from_tensor(tensor))
|
||||
}
|
||||
|
||||
/// Check convergence based on loss history
|
||||
fn check_convergence(&self, loss_history: &VecDeque<f64>, current_loss: f64) -> Result<bool> {
|
||||
if loss_history.len() < self.config.patience / 2 {
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
// Check if loss has plateaued
|
||||
let recent_losses: Vec<f64> = loss_history.iter().rev().take(10).copied().collect();
|
||||
let _mean_recent = recent_losses.iter().sum::<f64>() / recent_losses.len() as f64;
|
||||
|
||||
// Check relative improvement
|
||||
let relative_improvement = (recent_losses[0] - current_loss).abs() / recent_losses[0].abs();
|
||||
|
||||
Ok(current_loss < self.config.convergence_tolerance || relative_improvement < 1e-8)
|
||||
}
|
||||
|
||||
/// Update loss weights adaptively based on loss magnitudes
|
||||
async fn update_loss_weights(
|
||||
&self,
|
||||
_pinn: &PINN,
|
||||
loss_components: &LossComponents,
|
||||
) -> Result<()> {
|
||||
// Adaptive weight balancing to prevent one loss from dominating
|
||||
let physics_mag = loss_components.physics.log10().abs();
|
||||
let boundary_mag = loss_components.boundary.log10().abs();
|
||||
|
||||
// This would update the weights in self.config if it were mutable
|
||||
// For now, just log the recommended adjustments
|
||||
tracing::debug!(
|
||||
"Recommended weight adjustments - Physics: {:.3}, Boundary: {:.3}",
|
||||
1.0 / physics_mag.max(1e-6),
|
||||
1.0 / boundary_mag.max(1e-6)
|
||||
);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Convert coordinate points to tensor
|
||||
fn points_to_tensor(&self, points: &[(f64, f64)]) -> Result<Tensor> {
|
||||
let mut data = Vec::with_capacity(points.len() * 2);
|
||||
for (x, t) in points {
|
||||
data.push(*x as f32);
|
||||
data.push(*t as f32);
|
||||
}
|
||||
|
||||
Tensor::from_slice(&data, &[points.len(), 2], &Device::Cuda(0))
|
||||
.map_err(|e| ScienceError::computation(format!("Failed to create tensor: {e}")))
|
||||
}
|
||||
}
|
||||
|
||||
/// Loss weighting strategies
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub enum LossWeighting {
|
||||
/// Fixed weights
|
||||
Fixed,
|
||||
/// Adaptive weights based on loss magnitudes
|
||||
Adaptive,
|
||||
/// Curriculum learning
|
||||
Curriculum,
|
||||
}
|
||||
|
||||
/// Scheduler configuration
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct SchedulerConfig {
|
||||
/// Initial learning rate
|
||||
pub initial_lr: f64,
|
||||
/// Decay factor
|
||||
pub decay_factor: f64,
|
||||
/// Decay steps
|
||||
pub decay_steps: usize,
|
||||
}
|
||||
|
||||
/// Adam optimizer implementation for PINN training
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct AdamOptimizer {
|
||||
learning_rate: f64,
|
||||
beta1: f64,
|
||||
beta2: f64,
|
||||
epsilon: f64,
|
||||
iteration: usize,
|
||||
// These would store momentum and velocity for each parameter
|
||||
// In a real implementation, these would be proper tensor storage
|
||||
}
|
||||
|
||||
impl AdamOptimizer {
|
||||
#[must_use]
|
||||
pub fn new(learning_rate: f64) -> Self {
|
||||
Self {
|
||||
learning_rate,
|
||||
beta1: 0.9,
|
||||
beta2: 0.999,
|
||||
epsilon: 1e-8,
|
||||
iteration: 0,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn step(&mut self, _parameters: Vec<Variable>) -> Result<()> {
|
||||
self.iteration += 1;
|
||||
// Real implementation would update parameters using Adam algorithm
|
||||
// For now, this is a placeholder that maintains the API
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn set_learning_rate(&mut self, lr: f64) {
|
||||
self.learning_rate = lr;
|
||||
}
|
||||
}
|
||||
|
||||
/// Cosine annealing learning rate scheduler
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct CosineAnnealingScheduler {
|
||||
initial_lr: f64,
|
||||
final_lr: f64,
|
||||
total_steps: usize,
|
||||
}
|
||||
|
||||
impl CosineAnnealingScheduler {
|
||||
#[must_use]
|
||||
pub fn new(initial_lr: f64, total_steps: usize, final_lr: f64) -> Self {
|
||||
Self {
|
||||
initial_lr,
|
||||
final_lr,
|
||||
total_steps,
|
||||
}
|
||||
}
|
||||
|
||||
#[must_use]
|
||||
pub fn step(&self, current_step: usize) -> f64 {
|
||||
let progress = current_step as f64 / self.total_steps as f64;
|
||||
let cosine_decay = 0.5 * (1.0 + (std::f64::consts::PI * progress).cos());
|
||||
self.final_lr + (self.initial_lr - self.final_lr) * cosine_decay
|
||||
}
|
||||
}
|
||||
|
||||
/// Physics gradients computed during forward pass
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct PhysicsGradients {
|
||||
pub du_dx: rtx_autograd::Variable, // ∂u/∂x
|
||||
pub du_dt: rtx_autograd::Variable, // ∂u/∂t
|
||||
pub d2u_dx2: rtx_autograd::Variable, // ∂²u/∂x²
|
||||
pub d2u_dt2: rtx_autograd::Variable, // ∂²u/∂t²
|
||||
pub d2u_dxdt: rtx_autograd::Variable, // ∂²u/∂x∂t
|
||||
}
|
||||
Reference in New Issue
Block a user