//! Inverse Helmholtz equation residual computation //! //! The Helmholtz equation for MRE is: //! //! ```text //! div(mu * grad(u)) + rho * omega^2 * u = 0 //! ``` //! //! Expanding the divergence: //! //! ```text //! R = mu * (u_xx + u_yy) + mu_x * u_x + mu_y * u_y + rho * omega^2 * u = 0 //! ``` //! //! where: //! - u is the complex wave displacement (from Wave Net) //! - mu is the shear modulus (from Stiffness Texture) //! - rho is tissue density //! - omega is angular frequency use crate::config::MreConfig; use crate::wave_net::WaveDerivatives; use anyhow::Result; use rtx_tensor::{Device, Tensor}; /// Helmholtz equation residual computer pub struct HelmholtzResidual { /// Non-dimensional rho * omega^2 coefficient rho_omega_sq: f32, /// Device device: Device, } impl HelmholtzResidual { /// Create a new Helmholtz residual computer pub fn new(config: &MreConfig, device: &Device) -> Self { Self { rho_omega_sq: config.nondim_rho_omega_sq(), device: device.clone(), } } /// Compute physics residual at collocation points /// /// R = mu * laplacian(u) + grad(mu) . grad(u) + rho*omega^2 * u /// /// # Arguments /// * `wave_derivs` - Wave field and all spatial derivatives from Wave Net /// * `mu` - Stiffness values at sample points [batch] /// * `mu_x` - dmu/dx at sample points [batch] /// * `mu_y` - dmu/dy at sample points [batch] /// /// # Returns /// (r_real, r_imag) - Residual for real and imaginary parts pub fn compute_residual( &self, wave_derivs: &WaveDerivatives, mu: &Tensor, mu_x: &Tensor, mu_y: &Tensor, ) -> Result<(Tensor, Tensor)> { // Real part residual let r_real = self.compute_residual_component( &wave_derivs.u_real, &wave_derivs.u_x_real, &wave_derivs.u_y_real, &wave_derivs.u_xx_real, &wave_derivs.u_yy_real, mu, mu_x, mu_y, )?; // Imaginary part residual let r_imag = self.compute_residual_component( &wave_derivs.u_imag, &wave_derivs.u_x_imag, &wave_derivs.u_y_imag, &wave_derivs.u_xx_imag, &wave_derivs.u_yy_imag, mu, mu_x, mu_y, )?; Ok((r_real, r_imag)) } /// Compute residual for a single component (real or imag) /// /// R = mu * (u_xx + u_yy) + mu_x * u_x + mu_y * u_y + rho*omega^2 * u fn compute_residual_component( &self, u: &Tensor, u_x: &Tensor, u_y: &Tensor, u_xx: &Tensor, u_yy: &Tensor, mu: &Tensor, mu_x: &Tensor, mu_y: &Tensor, ) -> Result { // Term 1: mu * laplacian(u) = mu * (u_xx + u_yy) let laplacian = u_xx.add(u_yy)?; let term1 = mu.mul(&laplacian)?; // Term 2: grad(mu) . grad(u) = mu_x * u_x + mu_y * u_y let mu_x_ux = mu_x.mul(u_x)?; let mu_y_uy = mu_y.mul(u_y)?; let term2 = mu_x_ux.add(&mu_y_uy)?; // Term 3: rho * omega^2 * u let term3 = u.mul_scalar(self.rho_omega_sq)?; // Total residual R = term1 + term2 + term3 let residual = term1.add(&term2)?.add(&term3)?; Ok(residual) } /// Compute physics loss: L_physics = mean(R_real^2 + R_imag^2) pub fn physics_loss(&self, r_real: &Tensor, r_imag: &Tensor) -> Result { let r_real_sq = r_real.square()?; let r_imag_sq = r_imag.square()?; let r_sq = r_real_sq.add(&r_imag_sq)?; // Mean over batch let sum = r_sq.sum(None)?; let sum_val = sum.to_cpu()?[0]; let batch_size = r_real.shape().dims()[0] as f32; Ok(sum_val / batch_size) } /// Compute gradient of loss w.r.t stiffness mu /// /// L = |R|^2 = R_real^2 + R_imag^2 /// dL/dmu = 2 * R_real * dR_real/dmu + 2 * R_imag * dR_imag/dmu /// /// From R = mu * laplacian + grad_mu . grad_u + rho*omega^2*u /// dR/dmu = laplacian(u) (ignoring neighbor dependency of grad_mu) /// /// This is a simplification (lagged diffusivity) that is standard in /// elastography inverse problems. pub fn compute_stiffness_gradient( &self, wave_derivs: &WaveDerivatives, r_real: &Tensor, r_imag: &Tensor, ) -> Result { // dR/dmu = laplacian(u) let lap_real = wave_derivs.laplacian_real()?; let lap_imag = wave_derivs.laplacian_imag()?; // dL/dmu = 2 * (R_real * lap_real + R_imag * lap_imag) let grad_real = r_real.mul(&lap_real)?; let grad_imag = r_imag.mul(&lap_imag)?; let grad = grad_real.add(&grad_imag)?.mul_scalar(2.0)?; Ok(grad) } } /// Compute data loss between predicted and measured wave fields pub fn compute_data_loss( pred_real: &Tensor, pred_imag: &Tensor, meas_real: &Tensor, meas_imag: &Tensor, ) -> Result { // L_data = mean(|u_pred - u_meas|^2) let diff_real = pred_real.sub(meas_real)?; let diff_imag = pred_imag.sub(meas_imag)?; let diff_real_sq = diff_real.square()?; let diff_imag_sq = diff_imag.square()?; let diff_sq = diff_real_sq.add(&diff_imag_sq)?; let sum = diff_sq.sum(None)?; let sum_val = sum.to_cpu()?[0]; let batch_size = pred_real.shape().dims()[0] as f32; Ok(sum_val / batch_size) } #[cfg(test)] mod tests { use super::*; fn get_test_device() -> Device { Device::try_default().unwrap() } fn get_test_config() -> MreConfig { MreConfig::fast() } #[test] fn test_residual_creation() { let config = get_test_config(); let device = get_test_device(); let residual = HelmholtzResidual::new(&config, &device); assert!(residual.rho_omega_sq > 0.0); } #[test] fn test_zero_residual_for_constant_wave() { let config = get_test_config(); let device = get_test_device(); let residual = HelmholtzResidual::new(&config, &device); // For constant wave field (u = c), all derivatives are zero // R = mu * 0 + mu_x * 0 + mu_y * 0 + rho*omega^2 * c // R should equal rho*omega^2 * c (not zero unless c=0) let batch = 10; let u = Tensor::from_data(vec![0.0; batch], vec![batch], &device).unwrap(); let u_x = Tensor::from_data(vec![0.0; batch], vec![batch], &device).unwrap(); let u_y = Tensor::from_data(vec![0.0; batch], vec![batch], &device).unwrap(); let u_xx = Tensor::from_data(vec![0.0; batch], vec![batch], &device).unwrap(); let u_yy = Tensor::from_data(vec![0.0; batch], vec![batch], &device).unwrap(); let wave_derivs = WaveDerivatives { u_real: u.clone(), u_imag: u.clone(), u_x_real: u_x.clone(), u_x_imag: u_x.clone(), u_y_real: u_y.clone(), u_y_imag: u_y.clone(), u_xx_real: u_xx.clone(), u_xx_imag: u_xx.clone(), u_yy_real: u_yy.clone(), u_yy_imag: u_yy, }; let mu = Tensor::from_data(vec![1.0; batch], vec![batch], &device).unwrap(); let mu_x = Tensor::from_data(vec![0.0; batch], vec![batch], &device).unwrap(); let mu_y = Tensor::from_data(vec![0.0; batch], vec![batch], &device).unwrap(); let (r_real, r_imag) = residual .compute_residual(&wave_derivs, &mu, &mu_x, &mu_y) .unwrap(); // For u=0 everywhere, residual should be zero let loss = residual.physics_loss(&r_real, &r_imag).unwrap(); assert!(loss.abs() < 1e-10); } #[test] fn test_data_loss() { let device = get_test_device(); let pred_real = Tensor::from_data(vec![1.0, 2.0, 3.0], vec![3], &device).unwrap(); let pred_imag = Tensor::from_data(vec![0.0, 0.0, 0.0], vec![3], &device).unwrap(); let meas_real = Tensor::from_data(vec![1.0, 2.0, 3.0], vec![3], &device).unwrap(); let meas_imag = Tensor::from_data(vec![0.0, 0.0, 0.0], vec![3], &device).unwrap(); let loss = compute_data_loss(&pred_real, &pred_imag, &meas_real, &meas_imag).unwrap(); // Perfect match should give zero loss assert!(loss.abs() < 1e-10); } }