//! Main bioheat solver combining all components use crate::{ config::BioheatConfig, network::{Coordinate4D, ThermalPinn}, pennes::{PennesResidual, TemperatureDerivatives}, probe::ProbeHeatSource, tissue::TissueDomain, }; use bioheat_shared::{ AblationZone, BioheatLossRecord, BioheatSnapshot, BioheatStatus, Point3D, ProbeGeometry, TemperatureField, }; use rand::SeedableRng; use rand_chacha::ChaCha8Rng; use std::time::Instant; /// AdamW optimizer for neural network training #[derive(Debug, Clone)] struct AdamW { /// First moment estimates (momentum) for weights m_weights: Vec>>, /// Second moment estimates (variance) for weights v_weights: Vec>>, /// First moment estimates for biases m_biases: Vec>, /// Second moment estimates for biases v_biases: Vec>, /// First moment estimates for Fourier features m_fourier: Vec>, /// Second moment estimates for Fourier features v_fourier: Vec>, /// Timestep counter for bias correction t: u64, /// Learning rate lr: f64, /// First moment decay rate (β₁) beta1: f64, /// Second moment decay rate (β₂) beta2: f64, /// Numerical stability constant (ε) eps: f64, /// Weight decay coefficient (λ) weight_decay: f64, } impl AdamW { /// Create a new AdamW optimizer matching the network architecture fn new( weights_shape: &[Vec>], biases_shape: &[Vec], fourier_shape: &[Vec], lr: f64, ) -> Self { // Initialize momentum and variance to zero let m_weights = weights_shape .iter() .map(|layer| layer.iter().map(|row| vec![0.0; row.len()]).collect()) .collect(); let v_weights = weights_shape .iter() .map(|layer| layer.iter().map(|row| vec![0.0; row.len()]).collect()) .collect(); let m_biases = biases_shape.iter().map(|b| vec![0.0; b.len()]).collect(); let v_biases = biases_shape.iter().map(|b| vec![0.0; b.len()]).collect(); let m_fourier = fourier_shape .iter() .map(|row| vec![0.0; row.len()]) .collect(); let v_fourier = fourier_shape .iter() .map(|row| vec![0.0; row.len()]) .collect(); Self { m_weights, v_weights, m_biases, v_biases, m_fourier, v_fourier, t: 0, lr, beta1: 0.9, beta2: 0.999, eps: 1e-8, weight_decay: 0.01, } } /// Perform one optimization step fn step(&mut self, grads: &NetworkGradients, pinn: &mut ThermalPinn) { self.t += 1; // Bias correction factors let bias_correction1 = 1.0 - self.beta1.powi(self.t as i32); let bias_correction2 = 1.0 - self.beta2.powi(self.t as i32); // Update weights let weights = pinn.weights_mut(); for (layer_idx, (layer_grads, layer_weights)) in grads.weights.iter().zip(weights.iter_mut()).enumerate() { for (row_idx, (grad_row, weight_row)) in layer_grads.iter().zip(layer_weights.iter_mut()).enumerate() { for (col_idx, (grad, weight)) in grad_row.iter().zip(weight_row.iter_mut()).enumerate() { // Update first moment: m = β₁ * m + (1 - β₁) * g self.m_weights[layer_idx][row_idx][col_idx] = self.beta1 * self.m_weights[layer_idx][row_idx][col_idx] + (1.0 - self.beta1) * grad; // Update second moment: v = β₂ * v + (1 - β₂) * g² self.v_weights[layer_idx][row_idx][col_idx] = self.beta2 * self.v_weights[layer_idx][row_idx][col_idx] + (1.0 - self.beta2) * grad * grad; // Bias-corrected estimates let m_hat = self.m_weights[layer_idx][row_idx][col_idx] / bias_correction1; let v_hat = self.v_weights[layer_idx][row_idx][col_idx] / bias_correction2; // AdamW update with decoupled weight decay let adam_update = m_hat / (v_hat.sqrt() + self.eps); *weight -= self.lr * (adam_update + self.weight_decay * *weight); } } } // Update biases let biases = pinn.biases_mut(); for (layer_idx, (layer_grads, layer_biases)) in grads.biases.iter().zip(biases.iter_mut()).enumerate() { for (idx, (grad, bias)) in layer_grads.iter().zip(layer_biases.iter_mut()).enumerate() { self.m_biases[layer_idx][idx] = self.beta1 * self.m_biases[layer_idx][idx] + (1.0 - self.beta1) * grad; self.v_biases[layer_idx][idx] = self.beta2 * self.v_biases[layer_idx][idx] + (1.0 - self.beta2) * grad * grad; let m_hat = self.m_biases[layer_idx][idx] / bias_correction1; let v_hat = self.v_biases[layer_idx][idx] / bias_correction2; let adam_update = m_hat / (v_hat.sqrt() + self.eps); *bias -= self.lr * adam_update; // No weight decay on biases } } // Update Fourier features let fourier_b = pinn.fourier_b_mut(); for (row_idx, (grad_row, fourier_row)) in grads.fourier_b.iter().zip(fourier_b.iter_mut()).enumerate() { for (col_idx, (grad, fourier_val)) in grad_row.iter().zip(fourier_row.iter_mut()).enumerate() { self.m_fourier[row_idx][col_idx] = self.beta1 * self.m_fourier[row_idx][col_idx] + (1.0 - self.beta1) * grad; self.v_fourier[row_idx][col_idx] = self.beta2 * self.v_fourier[row_idx][col_idx] + (1.0 - self.beta2) * grad * grad; let m_hat = self.m_fourier[row_idx][col_idx] / bias_correction1; let v_hat = self.v_fourier[row_idx][col_idx] / bias_correction2; let adam_update = m_hat / (v_hat.sqrt() + self.eps); *fourier_val -= self.lr * (adam_update + self.weight_decay * *fourier_val); } } } } /// Gradients for all network parameters #[derive(Debug, Clone)] struct NetworkGradients { /// Gradients for weight matrices weights: Vec>>, /// Gradients for bias vectors biases: Vec>, /// Gradients for Fourier feature matrix fourier_b: Vec>, } impl NetworkGradients { /// Create zero gradients matching network architecture fn zeros_like(pinn: &ThermalPinn) -> Self { // Clone the network to access shapes let pinn_clone = pinn.clone(); let mut temp_pinn = pinn_clone; // Get shapes by accessing each mutable reference separately let grad_weights = { let weights = temp_pinn.weights_mut(); weights .iter() .map(|layer| layer.iter().map(|row| vec![0.0; row.len()]).collect()) .collect() }; let grad_biases = { let biases = temp_pinn.biases_mut(); biases.iter().map(|b| vec![0.0; b.len()]).collect() }; let grad_fourier = { let fourier_b = temp_pinn.fourier_b_mut(); fourier_b.iter().map(|row| vec![0.0; row.len()]).collect() }; Self { weights: grad_weights, biases: grad_biases, fourier_b: grad_fourier, } } } /// Main bioheat PINN solver pub struct BioheatSolver { /// Configuration config: BioheatConfig, /// Neural network for temperature prediction pinn: ThermalPinn, /// Pennes equation residual computer pennes: PennesResidual, /// Heat source from ablation probe probe_source: ProbeHeatSource, /// Tissue domain domain: TissueDomain, /// Current training step step: usize, /// Current simulation time current_time: f32, /// Training loss history loss_history: Vec, /// Random number generator (seeded for reproducibility) rng: ChaCha8Rng, /// Training start time (for throughput calculation) train_start: Option, /// AdamW optimizer optimizer: Option, } impl BioheatSolver { /// Create a new bioheat solver with given configuration #[must_use] pub fn new(config: BioheatConfig) -> Self { let pinn = ThermalPinn::new(config.network.clone()); let pennes = PennesResidual::new(config.physics.clone()); let probe_source = ProbeHeatSource::new(config.probe.clone(), config.probe_power); let domain = TissueDomain::new(config.domain, (config.time.t_start, config.time.t_end)); Self { config, pinn, pennes, probe_source, domain, step: 0, current_time: 0.0, loss_history: Vec::new(), rng: ChaCha8Rng::seed_from_u64(42), train_start: None, optimizer: None, } } /// Create solver with default liver ablation configuration #[must_use] pub fn liver_default() -> Self { Self::new(BioheatConfig::liver_default()) } /// Perform a single training step pub fn step(&mut self) -> BioheatLossRecord { if self.train_start.is_none() { self.train_start = Some(Instant::now()); } // Initialize optimizer on first step if self.optimizer.is_none() { // Create optimizer by cloning shapes let pinn_clone = self.pinn.clone(); let mut temp_pinn = pinn_clone; let weights_shape = { let w = temp_pinn.weights_mut(); w.clone() }; let biases_shape = { let b = temp_pinn.biases_mut(); b.clone() }; let fourier_shape = { let f = temp_pinn.fourier_b_mut(); f.clone() }; self.optimizer = Some(AdamW::new( &weights_shape, &biases_shape, &fourier_shape, self.config.training.learning_rate as f64, )); } // Sample training points let collocation_points = self .domain .sample_interior_batch(&mut self.rng, self.config.training.num_collocation); let boundary_points = self .domain .sample_boundary_batch(&mut self.rng, self.config.training.num_boundary); let initial_points = self .domain .sample_initial_batch(&mut self.rng, self.config.training.num_initial); // Compute losses let physics_loss = self.compute_physics_loss(&collocation_points); let boundary_loss = self.compute_boundary_loss(&boundary_points); let initial_loss = self.compute_initial_loss(&initial_points); let probe_loss = self.compute_probe_loss(&collocation_points); // Weight and combine losses let weights = &self.config.training.weights; let _total_loss = weights.physics * physics_loss + weights.boundary * boundary_loss + weights.initial * initial_loss + weights.probe * probe_loss; // Compute gradients using finite differences let grads = self.compute_network_gradients(); // Update weights using AdamW if let Some(optimizer) = &mut self.optimizer { optimizer.step(&grads, &mut self.pinn); } self.step += 1; let record = BioheatLossRecord::new( self.step, physics_loss, boundary_loss, initial_loss, probe_loss, ); self.loss_history.push(record.clone()); record } /// Train for multiple steps pub fn train(&mut self, num_steps: usize) -> Vec { let mut losses = Vec::with_capacity(num_steps); for _ in 0..num_steps { losses.push(self.step()); } losses } /// Compute physics (PDE residual) loss fn compute_physics_loss(&self, points: &[(Point3D, f32)]) -> f32 { let eps = 1e-4; // For finite difference derivatives let mut total_residual_sq = 0.0; for (point, t) in points { let coord = Coordinate4D::from_point_and_time(*point, *t); // Get temperature and derivatives let temp = self.pinn.forward(&coord); let (dt_dx, dt_dy, dt_dz, dt_dt) = self.pinn.gradients(&coord, eps); let (d2t_dx2, d2t_dy2, d2t_dz2) = self.pinn.second_derivatives(&coord, eps); let derivs = TemperatureDerivatives { t: temp as f32, dt_dt: dt_dt as f32, dt_dx: dt_dx as f32, dt_dy: dt_dy as f32, dt_dz: dt_dz as f32, d2t_dx2: d2t_dx2 as f32, d2t_dy2: d2t_dy2 as f32, d2t_dz2: d2t_dz2 as f32, }; // Get heat source at this point let heat_source = self.probe_source.heat_source_at(point); // Compute residual let residual = self.pennes.compute(&derivs, heat_source); total_residual_sq += residual * residual; } (total_residual_sq / points.len() as f32).sqrt() } /// Compute boundary condition loss (Dirichlet: T = T_body at boundary) fn compute_boundary_loss(&self, points: &[(Point3D, f32)]) -> f32 { let target_temp = self.config.physics.body_temperature; let mut total_error_sq = 0.0; for (point, t) in points { let coord = Coordinate4D::from_point_and_time(*point, *t); let predicted_temp = self.pinn.forward(&coord) as f32; let error = predicted_temp - target_temp; total_error_sq += error * error; } (total_error_sq / points.len() as f32).sqrt() } /// Compute initial condition loss (T = T_body at t=0) fn compute_initial_loss(&self, points: &[(Point3D, f32)]) -> f32 { let target_temp = self.config.physics.body_temperature; let mut total_error_sq = 0.0; for (point, t) in points { let coord = Coordinate4D::from_point_and_time(*point, *t); let predicted_temp = self.pinn.forward(&coord) as f32; let error = predicted_temp - target_temp; total_error_sq += error * error; } (total_error_sq / points.len() as f32).sqrt() } /// Compute probe heat source condition loss fn compute_probe_loss(&self, points: &[(Point3D, f32)]) -> f32 { // For points near the probe, temperature should be elevated // This is a soft constraint to help convergence let mut total_loss = 0.0; let mut count = 0; for (point, t) in points { let heat_source = self.probe_source.heat_source_at(point); if heat_source > 1e3 { // Near probe (significant heat) let coord = Coordinate4D::from_point_and_time(*point, *t); let predicted_temp = self.pinn.forward(&coord) as f32; // Temperature should be above body temp when probe is active if predicted_temp < self.config.physics.body_temperature { let error = self.config.physics.body_temperature - predicted_temp; total_loss += error * error; } count += 1; } } if count > 0 { (total_loss / count as f32).sqrt() } else { 0.0 } } /// Generate a snapshot of the current state for visualization #[must_use] pub fn snapshot(&self, resolution: (usize, usize, usize)) -> BioheatSnapshot { let grid_points = self.domain.regular_grid(resolution); let mut values = Vec::with_capacity(grid_points.len()); // Evaluate temperature at each grid point for current time for point in &grid_points { let coord = Coordinate4D::from_point_and_time(*point, self.current_time); values.push(self.pinn.forward(&coord) as f32); } let temperature = TemperatureField { resolution, values, bounds: self.domain.bounds, }; let max_temperature = temperature.max_temperature(); let ablation_volume = temperature.ablated_volume(); // Create ablation zone (simplified - no mesh generation) let ablation_zone = AblationZone::from_volume(ablation_volume, 60.0); let loss = self.loss_history.last().cloned().unwrap_or_default(); BioheatSnapshot { temperature, ablation_zone, probe: self.probe_source.geometry.clone(), step: self.step, time: self.current_time, loss, max_temperature, } } /// Get current solver status #[must_use] pub fn status(&self) -> BioheatStatus { let steps_per_second = if let Some(start) = self.train_start { let elapsed = start.elapsed().as_secs_f32(); if elapsed > 0.0 { self.step as f32 / elapsed } else { 0.0 } } else { 0.0 }; // Quick estimate of max temp and ablation volume let center_coord = Coordinate4D::from_point_and_time(Point3D::origin(), self.current_time); let max_temperature = self.pinn.forward(¢er_coord) as f32; BioheatStatus { initialized: true, training: self.train_start.is_some(), step: self.step, total_steps: self.step, simulation_time: self.current_time, steps_per_second, max_temperature, ablation_volume_mm3: 0.0, // Would need full grid evaluation } } /// Advance simulation time pub fn advance_time(&mut self, dt: f32) { self.current_time = (self.current_time + dt).min(self.config.time.t_end); } /// Set simulation time directly pub fn set_time(&mut self, t: f32) { self.current_time = t.clamp(self.config.time.t_start, self.config.time.t_end); } /// Update probe position pub fn update_probe_position(&mut self, position: Point3D) { self.probe_source.geometry.position = position; self.config.probe.position = position; } /// Update probe power pub fn update_probe_power(&mut self, power: f32) { self.probe_source.set_power(power); self.config.probe_power = power; } /// Update probe geometry pub fn update_probe(&mut self, probe: ProbeGeometry) { self.probe_source.geometry = probe.clone(); self.config.probe = probe; } /// Compute gradients of total loss with respect to network parameters using finite differences fn compute_network_gradients(&self) -> NetworkGradients { // Use finite difference method to compute gradients const EPS: f64 = 1e-4; // Initialize gradients let mut grads = NetworkGradients::zeros_like(&self.pinn); // Sample a batch of points for gradient computation let mut rng = ChaCha8Rng::seed_from_u64(self.step as u64 + 42); let collocation_points = self .domain .sample_interior_batch(&mut rng, self.config.training.num_collocation); let boundary_points = self .domain .sample_boundary_batch(&mut rng, self.config.training.num_boundary); let initial_points = self .domain .sample_initial_batch(&mut rng, self.config.training.num_initial); // Compute base loss (for reference, not used in gradient computation) let _base_loss = self.compute_total_loss(&collocation_points, &boundary_points, &initial_points); // Clone network for perturbation let mut pinn_perturbed = self.pinn.clone(); // Compute gradients for weights using stochastic parameter sampling // For efficiency, we sample a subset of parameters per batch let total_weight_params: usize = { let weights = pinn_perturbed.weights_mut(); weights .iter() .map(|layer| layer.iter().map(std::vec::Vec::len).sum::()) .sum() }; // Sample up to 256 parameters per step for gradient computation let max_samples = 256.min(total_weight_params); let sample_scale = total_weight_params as f64 / max_samples as f64; for sample_idx in 0..max_samples { // Deterministic sampling using step counter let flat_idx = (self.step * 31 + sample_idx * 17) % total_weight_params; // Convert flat index to (layer, row, col) let (layer_idx, row_idx, col_idx) = { let weights = pinn_perturbed.weights_mut(); let mut remaining = flat_idx; let mut layer_idx = 0; while layer_idx < weights.len() { let layer_size: usize = weights[layer_idx].iter().map(std::vec::Vec::len).sum(); if remaining < layer_size { break; } remaining -= layer_size; layer_idx += 1; } if layer_idx >= weights.len() { continue; } // Find row and column within layer let mut row_idx = 0; while row_idx < weights[layer_idx].len() && remaining >= weights[layer_idx][row_idx].len() { remaining -= weights[layer_idx][row_idx].len(); row_idx += 1; } if row_idx >= weights[layer_idx].len() { continue; } let col_idx = remaining; (layer_idx, row_idx, col_idx) }; // Compute gradient using central differences let original_val = { let weights = pinn_perturbed.weights_mut(); weights[layer_idx][row_idx][col_idx] }; // Perturb + { let weights = pinn_perturbed.weights_mut(); weights[layer_idx][row_idx][col_idx] = original_val + EPS; } let loss_plus = self.compute_total_loss_with_pinn( &pinn_perturbed, &collocation_points, &boundary_points, &initial_points, ); // Perturb - { let weights = pinn_perturbed.weights_mut(); weights[layer_idx][row_idx][col_idx] = original_val - EPS; } let loss_minus = self.compute_total_loss_with_pinn( &pinn_perturbed, &collocation_points, &boundary_points, &initial_points, ); // Restore original value { let weights = pinn_perturbed.weights_mut(); weights[layer_idx][row_idx][col_idx] = original_val; } // Compute gradient (scaled to account for sampling) let grad = ((loss_plus - loss_minus) / (2.0 * EPS)) * sample_scale; grads.weights[layer_idx][row_idx][col_idx] = grad; } // Compute gradients for biases (full gradient, they're small) let num_bias_layers = { let biases = pinn_perturbed.biases_mut(); biases.len() }; for layer_idx in 0..num_bias_layers { let num_biases = { let biases = pinn_perturbed.biases_mut(); biases[layer_idx].len() }; for bias_idx in 0..num_biases { let original_val = { let biases = pinn_perturbed.biases_mut(); biases[layer_idx][bias_idx] }; { let biases = pinn_perturbed.biases_mut(); biases[layer_idx][bias_idx] = original_val + EPS; } let loss_plus = self.compute_total_loss_with_pinn( &pinn_perturbed, &collocation_points, &boundary_points, &initial_points, ); { let biases = pinn_perturbed.biases_mut(); biases[layer_idx][bias_idx] = original_val - EPS; } let loss_minus = self.compute_total_loss_with_pinn( &pinn_perturbed, &collocation_points, &boundary_points, &initial_points, ); { let biases = pinn_perturbed.biases_mut(); biases[layer_idx][bias_idx] = original_val; } let grad = (loss_plus - loss_minus) / (2.0 * EPS); grads.biases[layer_idx][bias_idx] = grad; } } // Compute gradients for Fourier features (sample-based) let total_fourier_params: usize = { let fourier_b = pinn_perturbed.fourier_b_mut(); fourier_b.iter().map(std::vec::Vec::len).sum() }; let fourier_samples = 64.min(total_fourier_params); let fourier_scale = total_fourier_params as f64 / fourier_samples as f64; let fourier_cols = { let fourier_b = pinn_perturbed.fourier_b_mut(); if fourier_b.is_empty() { 0 } else { fourier_b[0].len() } }; for sample_idx in 0..fourier_samples { let flat_idx = (self.step * 37 + sample_idx * 23) % total_fourier_params; let row_idx = flat_idx / fourier_cols; let col_idx = flat_idx % fourier_cols; let fourier_rows = { let fourier_b = pinn_perturbed.fourier_b_mut(); fourier_b.len() }; if row_idx >= fourier_rows { continue; } let original_val = { let fourier_b = pinn_perturbed.fourier_b_mut(); fourier_b[row_idx][col_idx] }; { let fourier_b = pinn_perturbed.fourier_b_mut(); fourier_b[row_idx][col_idx] = original_val + EPS; } let loss_plus = self.compute_total_loss_with_pinn( &pinn_perturbed, &collocation_points, &boundary_points, &initial_points, ); { let fourier_b = pinn_perturbed.fourier_b_mut(); fourier_b[row_idx][col_idx] = original_val - EPS; } let loss_minus = self.compute_total_loss_with_pinn( &pinn_perturbed, &collocation_points, &boundary_points, &initial_points, ); { let fourier_b = pinn_perturbed.fourier_b_mut(); fourier_b[row_idx][col_idx] = original_val; } let grad = ((loss_plus - loss_minus) / (2.0 * EPS)) * fourier_scale; grads.fourier_b[row_idx][col_idx] = grad; } grads } /// Compute total weighted loss fn compute_total_loss( &self, collocation_points: &[(Point3D, f32)], boundary_points: &[(Point3D, f32)], initial_points: &[(Point3D, f32)], ) -> f64 { let physics_loss = self.compute_physics_loss(collocation_points); let boundary_loss = self.compute_boundary_loss(boundary_points); let initial_loss = self.compute_initial_loss(initial_points); let probe_loss = self.compute_probe_loss(collocation_points); let weights = &self.config.training.weights; (weights.physics * physics_loss + weights.boundary * boundary_loss + weights.initial * initial_loss + weights.probe * probe_loss) as f64 } /// Compute total weighted loss with a specific PINN instance fn compute_total_loss_with_pinn( &self, pinn: &ThermalPinn, collocation_points: &[(Point3D, f32)], boundary_points: &[(Point3D, f32)], initial_points: &[(Point3D, f32)], ) -> f64 { let physics_loss = self.compute_physics_loss_with_pinn(pinn, collocation_points); let boundary_loss = self.compute_boundary_loss_with_pinn(pinn, boundary_points); let initial_loss = self.compute_initial_loss_with_pinn(pinn, initial_points); let probe_loss = self.compute_probe_loss_with_pinn(pinn, collocation_points); let weights = &self.config.training.weights; (weights.physics * physics_loss + weights.boundary * boundary_loss + weights.initial * initial_loss + weights.probe * probe_loss) as f64 } /// Compute physics loss with a specific PINN instance fn compute_physics_loss_with_pinn(&self, pinn: &ThermalPinn, points: &[(Point3D, f32)]) -> f32 { let eps = 1e-4; let mut total_residual_sq = 0.0; for (point, t) in points { let coord = Coordinate4D::from_point_and_time(*point, *t); let temp = pinn.forward(&coord); let (dt_dx, dt_dy, dt_dz, dt_dt) = pinn.gradients(&coord, eps); let (d2t_dx2, d2t_dy2, d2t_dz2) = pinn.second_derivatives(&coord, eps); let derivs = TemperatureDerivatives { t: temp as f32, dt_dt: dt_dt as f32, dt_dx: dt_dx as f32, dt_dy: dt_dy as f32, dt_dz: dt_dz as f32, d2t_dx2: d2t_dx2 as f32, d2t_dy2: d2t_dy2 as f32, d2t_dz2: d2t_dz2 as f32, }; let heat_source = self.probe_source.heat_source_at(point); let residual = self.pennes.compute(&derivs, heat_source); total_residual_sq += residual * residual; } (total_residual_sq / points.len() as f32).sqrt() } /// Compute boundary loss with a specific PINN instance fn compute_boundary_loss_with_pinn( &self, pinn: &ThermalPinn, points: &[(Point3D, f32)], ) -> f32 { let target_temp = self.config.physics.body_temperature; let mut total_error_sq = 0.0; for (point, t) in points { let coord = Coordinate4D::from_point_and_time(*point, *t); let predicted_temp = pinn.forward(&coord) as f32; let error = predicted_temp - target_temp; total_error_sq += error * error; } (total_error_sq / points.len() as f32).sqrt() } /// Compute initial loss with a specific PINN instance fn compute_initial_loss_with_pinn(&self, pinn: &ThermalPinn, points: &[(Point3D, f32)]) -> f32 { let target_temp = self.config.physics.body_temperature; let mut total_error_sq = 0.0; for (point, t) in points { let coord = Coordinate4D::from_point_and_time(*point, *t); let predicted_temp = pinn.forward(&coord) as f32; let error = predicted_temp - target_temp; total_error_sq += error * error; } (total_error_sq / points.len() as f32).sqrt() } /// Compute probe loss with a specific PINN instance fn compute_probe_loss_with_pinn(&self, pinn: &ThermalPinn, points: &[(Point3D, f32)]) -> f32 { let mut total_loss = 0.0; let mut count = 0; for (point, t) in points { let heat_source = self.probe_source.heat_source_at(point); if heat_source > 1e3 { let coord = Coordinate4D::from_point_and_time(*point, *t); let predicted_temp = pinn.forward(&coord) as f32; if predicted_temp < self.config.physics.body_temperature { let error = self.config.physics.body_temperature - predicted_temp; total_loss += error * error; } count += 1; } } if count > 0 { (total_loss / count as f32).sqrt() } else { 0.0 } } /// Reset the solver to initial state pub fn reset(&mut self) { self.step = 0; self.current_time = self.config.time.t_start; self.loss_history.clear(); self.train_start = None; self.pinn = ThermalPinn::new(self.config.network.clone()); self.optimizer = None; } /// Get current step #[must_use] pub fn current_step(&self) -> usize { self.step } /// Get current time #[must_use] pub fn current_time(&self) -> f32 { self.current_time } /// Get loss history #[must_use] pub fn loss_history(&self) -> &[BioheatLossRecord] { &self.loss_history } /// Get the configuration #[must_use] pub fn config(&self) -> &BioheatConfig { &self.config } } #[cfg(test)] mod tests { use super::*; #[test] fn test_solver_creation() { let solver = BioheatSolver::liver_default(); assert_eq!(solver.current_step(), 0); } #[test] fn test_single_step() { let mut solver = BioheatSolver::new(BioheatConfig::benchmark()); let loss = solver.step(); assert_eq!(loss.step, 1); assert!(loss.total_loss >= 0.0); } #[test] fn test_train_multiple() { let mut solver = BioheatSolver::new(BioheatConfig::benchmark()); let losses = solver.train(10); assert_eq!(losses.len(), 10); assert_eq!(solver.current_step(), 10); } #[test] fn test_snapshot() { let solver = BioheatSolver::new(BioheatConfig::benchmark()); let snapshot = solver.snapshot((8, 8, 8)); assert_eq!(snapshot.temperature.len(), 512); } #[test] fn test_status() { let solver = BioheatSolver::liver_default(); let status = solver.status(); assert!(status.initialized); assert!(!status.training); } #[test] fn test_time_advancement() { let mut solver = BioheatSolver::liver_default(); solver.advance_time(10.0); assert!((solver.current_time() - 10.0).abs() < 1e-6); // Should clamp at t_end solver.advance_time(1000.0); assert!(solver.current_time() <= solver.config().time.t_end); } #[test] fn test_reset() { let mut solver = BioheatSolver::new(BioheatConfig::benchmark()); solver.train(5); solver.advance_time(50.0); solver.reset(); assert_eq!(solver.current_step(), 0); assert!((solver.current_time()).abs() < 1e-6); assert!(solver.loss_history().is_empty()); } #[test] fn test_probe_update() { let mut solver = BioheatSolver::liver_default(); let new_pos = Point3D::new(0.01, 0.02, 0.03); solver.update_probe_position(new_pos); assert!((solver.config().probe.position.x - 0.01).abs() < 1e-6); } // ============================================================================ // TDD Tests for Weight Updates (RED phase - these should FAIL initially) // ============================================================================ #[test] fn test_weights_change_after_training_step() { // Test that neural network weights actually change after a training step let mut solver = BioheatSolver::new(BioheatConfig::benchmark()); // Capture initial weight snapshot let initial_weights = capture_network_weights(&solver.pinn); // Perform a training step solver.step(); // Capture weights after training let updated_weights = capture_network_weights(&solver.pinn); // At least some weights should have changed let weights_changed = weights_are_different(&initial_weights, &updated_weights); assert!( weights_changed, "Neural network weights must change after a training step" ); } #[test] fn test_loss_decreases_over_training() { // Test that loss decreases over multiple training epochs let mut solver = BioheatSolver::new(BioheatConfig::benchmark()); // Train for several steps let losses = solver.train(20); // Loss should generally decrease (allowing for some noise) // Check that average loss in second half is less than first half let first_half_avg: f32 = losses.iter().take(10).map(|r| r.total_loss).sum::() / 10.0; let second_half_avg: f32 = losses.iter().skip(10).map(|r| r.total_loss).sum::() / 10.0; assert!( second_half_avg < first_half_avg, "Loss should decrease over training: first_half={}, second_half={}", first_half_avg, second_half_avg ); } #[test] fn test_gradient_computation_is_nonzero() { // Test that computed gradients are non-zero for typical loss values let solver = BioheatSolver::new(BioheatConfig::benchmark()); // Sample a few training points let mut rng = ChaCha8Rng::seed_from_u64(123); let points = solver.domain.sample_interior_batch(&mut rng, 5); // Compute loss let loss = solver.compute_physics_loss(&points); // Loss should be positive (network starts random) assert!( loss > 0.0, "Physics loss should be positive for untrained network" ); // Compute gradients let grads = solver.compute_network_gradients(); // At least some gradients should be non-zero let has_nonzero_grads = grads.weights.iter().any(|layer_weights: &Vec>| { layer_weights .iter() .any(|row: &Vec| row.iter().any(|&g: &f64| g.abs() > 1e-8)) }); assert!( has_nonzero_grads, "Gradients should be non-zero for non-zero loss" ); } #[test] fn test_optimizer_state_updates() { // Test that optimizer maintains state (momentum, variance) across steps let mut solver = BioheatSolver::new(BioheatConfig::benchmark()); // Initial optimizer state should exist after first step solver.step(); // Check optimizer has non-zero state (this will be implemented) assert!( solver.optimizer.is_some(), "Optimizer should be initialized after first training step" ); let opt = solver.optimizer.as_ref().unwrap(); // After first step, optimizer should have accumulated some state assert_eq!(opt.t, 1, "Optimizer timestep should be 1 after one step"); } #[test] fn test_learning_rate_is_applied() { // Test that different learning rates produce different weight updates let mut config1 = BioheatConfig::benchmark(); config1.training.learning_rate = 0.001; let mut config2 = BioheatConfig::benchmark(); config2.training.learning_rate = 0.01; // 10x larger let mut solver1 = BioheatSolver::new(config1); let mut solver2 = BioheatSolver::new(config2); // Set same random seed for reproducibility solver1.rng = ChaCha8Rng::seed_from_u64(42); solver2.rng = ChaCha8Rng::seed_from_u64(42); // Capture initial weights (should be same due to seed in new()) let initial1 = capture_network_weights(&solver1.pinn); let initial2 = capture_network_weights(&solver2.pinn); // Perform one training step solver1.step(); solver2.step(); // Capture updated weights let updated1 = capture_network_weights(&solver1.pinn); let updated2 = capture_network_weights(&solver2.pinn); // Compute weight change magnitudes let change1 = compute_weight_change_magnitude(&initial1, &updated1); let change2 = compute_weight_change_magnitude(&initial2, &updated2); // Higher learning rate should produce larger weight changes assert!( change2 > change1, "Higher learning rate should produce larger weight updates: lr=0.01 change={} vs lr=0.001 change={}", change2, change1 ); } // ============================================================================ // Helper functions for tests // ============================================================================ /// Capture a snapshot of network weights for comparison fn capture_network_weights(pinn: &ThermalPinn) -> NetworkWeightSnapshot { // Clone the network to access weights let mut pinn_clone = pinn.clone(); let weights = pinn_clone.weights_mut().clone(); let biases = pinn_clone.biases_mut().clone(); let fourier_b = pinn_clone.fourier_b_mut().clone(); NetworkWeightSnapshot { weights, biases, fourier_b, } } /// Check if two weight snapshots are different fn weights_are_different(w1: &NetworkWeightSnapshot, w2: &NetworkWeightSnapshot) -> bool { // Check if any weights have changed for (layer1, layer2) in w1.weights.iter().zip(&w2.weights) { for (row1, row2) in layer1.iter().zip(layer2) { for (val1, val2) in row1.iter().zip(row2) { if (val1 - val2).abs() > 1e-10 { return true; } } } } // Check biases for (b1, b2) in w1.biases.iter().zip(&w2.biases) { for (val1, val2) in b1.iter().zip(b2) { if (val1 - val2).abs() > 1e-10 { return true; } } } // Check Fourier features for (f1, f2) in w1.fourier_b.iter().zip(&w2.fourier_b) { for (val1, val2) in f1.iter().zip(f2) { if (val1 - val2).abs() > 1e-10 { return true; } } } false } /// Compute magnitude of weight change fn compute_weight_change_magnitude( before: &NetworkWeightSnapshot, after: &NetworkWeightSnapshot, ) -> f64 { let mut total_change = 0.0; // Sum squared differences for weights for (layer1, layer2) in before.weights.iter().zip(&after.weights) { for (row1, row2) in layer1.iter().zip(layer2) { for (val1, val2) in row1.iter().zip(row2) { let diff = val1 - val2; total_change += diff * diff; } } } // Sum squared differences for biases for (b1, b2) in before.biases.iter().zip(&after.biases) { for (val1, val2) in b1.iter().zip(b2) { let diff = val1 - val2; total_change += diff * diff; } } // Sum squared differences for Fourier features for (f1, f2) in before.fourier_b.iter().zip(&after.fourier_b) { for (val1, val2) in f1.iter().zip(f2) { let diff = val1 - val2; total_change += diff * diff; } } total_change.sqrt() } /// Snapshot of network weights for testing #[derive(Clone)] struct NetworkWeightSnapshot { weights: Vec>>, biases: Vec>, fourier_b: Vec>, } }