Files
rustytorch/demos/rtx-mre/src/helmholtz.rs
T
2026-03-04 00:08:42 +00:00

265 lines
8.2 KiB
Rust

//! 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<Tensor> {
// 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<f32> {
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<Tensor> {
// 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<f32> {
// 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);
}
}