//! Physics-Informed Denoising Diffusion Model (PIDDM). //! //! This module implements the core PIDDM model that combines diffusion //! models with physics-informed constraints for PDE solving. use rand::Rng; use rand::prelude::*; use thiserror::Error; use rtx_backend::Backend; use rtx_tensor::GenericTensor; use crate::scheduler::{DDPMScheduler, NoiseScheduler}; use crate::unet::DiffusionUNet; /// Errors that can occur in PIDDM operations. #[derive(Error, Debug)] pub enum PIDDMError { /// Invalid configuration #[error("Invalid configuration: {0}")] InvalidConfig(String), /// Shape mismatch #[error("Shape mismatch: expected {expected}, got {got}")] ShapeMismatch { /// The expected shape description expected: String, /// The actual shape received got: String, }, /// Training error #[error("Training error: {0}")] TrainingError(String), /// Sampling error #[error("Sampling error: {0}")] SamplingError(String), } /// Configuration for PIDDM model. #[derive(Debug, Clone)] pub struct PIDDMConfig { /// Weight for physics loss (λ in combined loss = L_diffusion + λ * L_physics) pub physics_weight: f32, /// Whether to use physics loss during sampling (physics guidance) pub physics_guided_sampling: bool, /// Guidance strength during sampling (if physics_guided_sampling is true) pub guidance_strength: f32, /// Whether to clamp predictions to [-1, 1] pub clamp_predictions: bool, } impl Default for PIDDMConfig { fn default() -> Self { Self { physics_weight: 0.1, physics_guided_sampling: false, guidance_strength: 0.0, clamp_predictions: true, } } } impl PIDDMConfig { /// Create config optimized for PDE solving. pub fn for_pde() -> Self { Self { physics_weight: 0.5, // Higher physics weight for PDE applications physics_guided_sampling: true, guidance_strength: 0.1, clamp_predictions: false, // PDEs may have unbounded solutions } } } /// Physics-Informed Denoising Diffusion Model. /// /// PIDDM combines diffusion models with physics constraints by: /// 1. Adding physics residual loss during training /// 2. Optionally using physics guidance during sampling /// /// # Training Loss /// /// The training loss combines diffusion loss and physics loss: /// ```text /// L = L_diffusion + λ * L_physics /// = ||ε - ε_θ(x_t, t)||² + λ * ||L[x̂_0] - f||² /// ``` /// /// where: /// - ε is the noise added to the clean sample /// - ε_θ is the predicted noise from the UNet /// - x̂_0 is the predicted clean sample /// - L[·] is the differential operator (e.g., Laplacian) /// - f is the forcing term pub struct PIDDM> { /// Noise scheduler scheduler: DDPMScheduler, /// UNet noise predictor unet: DiffusionUNet, /// Configuration config: PIDDMConfig, } impl> PIDDM { /// Create a new PIDDM model. pub fn new(scheduler: DDPMScheduler, unet: DiffusionUNet, config: PIDDMConfig) -> Self { Self { scheduler, unet, config, } } /// Get reference to the scheduler. pub fn scheduler(&self) -> &DDPMScheduler { &self.scheduler } /// Get reference to the UNet. pub fn unet(&self) -> &DiffusionUNet { &self.unet } /// Get reference to the configuration. pub fn config(&self) -> &PIDDMConfig { &self.config } /// Perform a training step. /// /// # Arguments /// * `x0` - Clean samples [batch, channels, height, width] /// * `physics_fn` - Function that computes physics residual from predicted x0 /// * `device` - Device /// /// # Returns /// Tuple of (total_loss, diffusion_loss, physics_loss) pub fn training_step( &self, x0: &GenericTensor, physics_fn: F, device: &B::Device, ) -> (f32, f32, f32) where F: Fn(&GenericTensor) -> GenericTensor, { let mut rng = thread_rng(); // Sample random timestep let t = rng.gen_range(0..self.scheduler.num_timesteps()); // Sample noise let noise = GenericTensor::::randn(x0.shape(), device); // Create noisy sample let xt = self.scheduler.add_noise(x0, &noise, t); // Predict noise let pred_noise = self.unet.forward(&xt, t, device); // Compute diffusion loss: ||ε - ε_θ||² let diff = noise.sub(&pred_noise); let diff_sq = diff.mul(&diff); let diffusion_loss = diff_sq.mean().to_vec()[0]; // Predict x0 from noisy sample and predicted noise let pred_x0 = self.scheduler.predict_x0(&xt, &pred_noise, t); // Optionally clamp predictions let pred_x0 = if self.config.clamp_predictions { clamp(&pred_x0, -1.0, 1.0) } else { pred_x0 }; // Compute physics loss: ||L[x̂_0] - f||² let physics_residual = physics_fn(&pred_x0); let physics_sq = physics_residual.mul(&physics_residual); let physics_loss = physics_sq.mean().to_vec()[0]; // Combined loss let total_loss = diffusion_loss + self.config.physics_weight * physics_loss; (total_loss, diffusion_loss, physics_loss) } /// Sample from the model using DDPM sampling. /// /// # Arguments /// * `shape` - Shape of samples to generate [batch, channels, height, width] /// * `device` - Device /// /// # Returns /// Generated samples pub fn sample(&self, shape: [usize; 4], device: &B::Device) -> GenericTensor { // Start from pure noise let mut x = GenericTensor::::randn(shape, device); // Reverse diffusion process for t in (0..self.scheduler.num_timesteps()).rev() { // Predict noise let pred_noise = self.unet.forward(&x, t, device); // Sample noise for stochastic step (not needed at t=0) let noise = if t > 0 { Some(GenericTensor::::randn(shape, device)) } else { None }; // Reverse step x = self.scheduler.step(&x, &pred_noise, t, noise.as_ref()); // Optional: clamp intermediate results if self.config.clamp_predictions { x = clamp(&x, -1.0, 1.0); } } x } /// Sample with physics guidance. /// /// During sampling, we modify the predicted x0 to reduce physics residual: /// ```text /// x̂_0 = x̂_0 - γ * ∇_{x̂_0} ||L[x̂_0] - f||² /// ``` /// /// # Arguments /// * `shape` - Shape of samples to generate /// * `physics_fn` - Function that computes physics residual /// * `physics_grad_fn` - Function that computes gradient of physics loss w.r.t. input /// * `device` - Device /// /// # Returns /// Generated samples satisfying physics constraints pub fn sample_with_guidance( &self, shape: [usize; 4], physics_fn: F, physics_grad_fn: G, device: &B::Device, ) -> GenericTensor where F: Fn(&GenericTensor) -> GenericTensor, G: Fn(&GenericTensor) -> GenericTensor, { let mut x = GenericTensor::::randn(shape, device); for t in (0..self.scheduler.num_timesteps()).rev() { // Predict noise let pred_noise = self.unet.forward(&x, t, device); // Predict x0 let mut pred_x0 = self.scheduler.predict_x0(&x, &pred_noise, t); // Apply physics guidance (gradient descent on physics loss) if self.config.physics_guided_sampling && self.config.guidance_strength > 0.0 { // Compute gradient of physics loss let grad = physics_grad_fn(&pred_x0); // Update x0 to reduce physics residual pred_x0 = pred_x0.sub(&grad.mul_scalar(self.config.guidance_strength)); } // Clamp if configured if self.config.clamp_predictions { pred_x0 = clamp(&pred_x0, -1.0, 1.0); } // Reconstruct noise from corrected x0 let alpha_cumprod = self.scheduler.alpha_cumprod(t); let corrected_noise = x .sub(&pred_x0.mul_scalar(alpha_cumprod.sqrt())) .mul_scalar(1.0 / (1.0 - alpha_cumprod).sqrt()); // Sample noise for stochastic step let noise = if t > 0 { Some(GenericTensor::::randn(shape, device)) } else { None }; // Reverse step with corrected noise prediction x = self.scheduler.step(&x, &corrected_noise, t, noise.as_ref()); } x } /// Compute the Laplacian of a 2D field using finite differences. /// /// This is a helper for common physics computations. /// ∇²u ≈ (u[i+1,j] + u[i-1,j] + u[i,j+1] + u[i,j-1] - 4*u[i,j]) / h² pub fn laplacian_2d(field: &GenericTensor, h: f32) -> GenericTensor { let shape = field.shape(); let [batch, channels, height, width] = [shape[0], shape[1], shape[2], shape[3]]; let field_vec = field.to_vec(); let mut laplacian = vec![0.0f32; batch * channels * height * width]; let h2 = h * h; for b in 0..batch { for c in 0..channels { for i in 1..height - 1 { for j in 1..width - 1 { let idx = |ii: usize, jj: usize| -> usize { b * channels * height * width + c * height * width + ii * width + jj }; let center = field_vec[idx(i, j)]; let left = field_vec[idx(i - 1, j)]; let right = field_vec[idx(i + 1, j)]; let down = field_vec[idx(i, j - 1)]; let up = field_vec[idx(i, j + 1)]; laplacian[idx(i, j)] = (left + right + down + up - 4.0 * center) / h2; } } } } GenericTensor::::from_slice(&laplacian, shape, &field.device()) } /// Create a Poisson residual function. /// /// For -∇²u = f, the residual is: R = -∇²u - f pub fn poisson_residual( u: &GenericTensor, f: &GenericTensor, h: f32, ) -> GenericTensor { let laplacian = Self::laplacian_2d(u, h); // Residual = -∇²u - f = -laplacian - f laplacian.mul_scalar(-1.0).sub(f) } } /// Clamp tensor values to [min, max]. fn clamp, const D: usize>( x: &GenericTensor, min: f32, max: f32, ) -> GenericTensor { let x_vec = x.to_vec(); let clamped: Vec = x_vec.iter().map(|&v| v.clamp(min, max)).collect(); GenericTensor::::from_slice(&clamped, x.shape(), &x.device()) } #[cfg(test)] mod tests { use super::*; use crate::unet::UNetConfig; use rtx_backend_cpu::{CpuBackend, CpuDevice}; fn create_test_model() -> PIDDM { let device = CpuDevice::default(); let scheduler = DDPMScheduler::new(100, 1e-4, 0.02); // Small for testing let unet_config = UNetConfig { in_channels: 1, out_channels: 1, base_channels: 16, hidden_mult: 2, time_embed_dim: 32, }; let spatial_size = 64; // 8x8 let unet = DiffusionUNet::new(unet_config, spatial_size, &device); let config = PIDDMConfig::default(); PIDDM::new(scheduler, unet, config) } #[test] fn test_piddm_creation() { let model = create_test_model(); assert_eq!(model.scheduler().num_timesteps(), 100); } #[test] fn test_laplacian_2d() { let device = CpuDevice::default(); // Create a simple quadratic field: u(x,y) = x² + y² // For this field, ∇²u = 2 + 2 = 4 (constant) let n = 8; let h = 1.0 / (n - 1) as f32; let mut data = vec![0.0f32; n * n]; for i in 0..n { for j in 0..n { let x = i as f32 * h; let y = j as f32 * h; data[i * n + j] = x * x + y * y; } } let field = GenericTensor::::from_slice(&data, [1, 1, n, n], &device); let laplacian = PIDDM::::laplacian_2d(&field, h); let laplacian_vec = laplacian.to_vec(); // Interior points should have Laplacian ≈ 4 let center = n / 2; let center_val = laplacian_vec[center * n + center]; assert!( (center_val - 4.0).abs() < 0.5, "Laplacian should be close to 4, got {}", center_val ); } #[test] fn test_training_step() { let model = create_test_model(); let device = CpuDevice::default(); // Create dummy clean sample let x0 = GenericTensor::::rand([1, 1, 8, 8], &device); // Dummy physics function (just returns zeros - no physics constraint) let physics_fn = |_x: &GenericTensor| { GenericTensor::::zeros([1, 1, 8, 8], &device) }; let (total, diffusion, physics) = model.training_step(&x0, physics_fn, &device); // Losses should be non-negative assert!(total >= 0.0, "Total loss should be non-negative"); assert!(diffusion >= 0.0, "Diffusion loss should be non-negative"); assert!(physics >= 0.0, "Physics loss should be non-negative"); // Physics loss should be 0 (since residual function returns zeros) assert!( physics < 1e-10, "Physics loss should be ~0 for zero residual" ); } #[test] fn test_sample_shape() { let model = create_test_model(); let device = CpuDevice::default(); // Note: This is a simplified test - full sampling is slow // We just test that the forward pass works with the expected shapes let shape = [1, 1, 8, 8]; // Test single step let x = GenericTensor::::randn(shape, &device); let pred_noise = model.unet().forward(&x, 50, &device); let pred_shape = pred_noise.shape(); assert_eq!(pred_shape, shape); } #[test] fn test_poisson_residual() { let device = CpuDevice::default(); let n = 8; let h = 1.0 / (n - 1) as f32; // Create u = sin(πx)sin(πy), which satisfies -∇²u = 2π²sin(πx)sin(πy) let pi = std::f32::consts::PI; let mut u_data = vec![0.0f32; n * n]; let mut f_data = vec![0.0f32; n * n]; for i in 0..n { for j in 0..n { let x = i as f32 * h; let y = j as f32 * h; u_data[i * n + j] = (pi * x).sin() * (pi * y).sin(); f_data[i * n + j] = 2.0 * pi * pi * (pi * x).sin() * (pi * y).sin(); } } let u = GenericTensor::::from_slice(&u_data, [1, 1, n, n], &device); let f = GenericTensor::::from_slice(&f_data, [1, 1, n, n], &device); let residual = PIDDM::::poisson_residual(&u, &f, h); let residual_vec = residual.to_vec(); // Interior residual should be small (discretization error) let center = n / 2; let center_residual = residual_vec[center * n + center]; assert!( center_residual.abs() < 1.0, "Poisson residual should be small at center, got {}", center_residual ); } }