Initial commit
This commit is contained in:
@@ -0,0 +1,264 @@
|
||||
//! 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);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user