Files
rustytorch/crates/specialized/rtx-neuro-fem/src/assembly.rs
T
2026-03-04 00:08:42 +00:00

675 lines
21 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! FEM stiffness matrix assembly for head modeling.
//!
//! Provides element-wise and global assembly of finite element
//! stiffness matrices using sparse matrix storage.
use crate::conductivity::{ConductivityTensor, TissueConductivity};
use crate::error::{FemError, FemResult};
use crate::mesh::{HeadMesh, TetElement};
use nalgebra::{Matrix4, Vector3};
use ndarray::Array2;
use rayon::prelude::*;
use serde::{Deserialize, Serialize};
use sprs::{CsMatI, TriMat};
/// Element stiffness matrix (4x4 for linear tetrahedron)
#[derive(Debug, Clone)]
pub struct ElementStiffness {
/// Element index
pub element_idx: usize,
/// Local stiffness matrix (4x4)
pub matrix: Matrix4<f64>,
/// Global node indices
pub node_indices: [usize; 4],
}
impl ElementStiffness {
/// Compute element stiffness matrix
///
/// For a linear tetrahedral element:
/// Ke[i,j] = ∫ (∇Ni)^T · σ · (∇Nj) dV
/// = V * (∇Ni)^T · σ · (∇Nj) (constant for linear elements)
pub fn compute(elem: &TetElement, mesh: &HeadMesh, conductivity: &ConductivityTensor) -> Self {
let grads = elem.shape_gradients(&mesh.nodes);
let sigma = &conductivity.tensor;
let vol = elem.volume;
let mut ke = Matrix4::zeros();
for i in 0..4 {
for j in 0..4 {
// Ke[i,j] = V * grad(Ni)^T * σ * grad(Nj)
let grad_i = &grads[i];
let grad_j = &grads[j];
ke[(i, j)] = vol * grad_i.dot(&(sigma * grad_j));
}
}
Self {
element_idx: elem.id,
matrix: ke,
node_indices: elem.nodes,
}
}
}
/// Global sparse stiffness matrix in CSR format
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StiffnessMatrix {
/// Number of rows (nodes)
pub n_rows: usize,
/// Number of columns (nodes)
pub n_cols: usize,
/// CSR row pointers
pub row_ptr: Vec<usize>,
/// CSR column indices
pub col_idx: Vec<usize>,
/// CSR values
pub values: Vec<f64>,
}
impl StiffnessMatrix {
/// Create from sprs CSR matrix
pub fn from_csr(csr: &CsMatI<f64, usize>) -> Self {
Self {
n_rows: csr.rows(),
n_cols: csr.cols(),
row_ptr: csr.indptr().as_slice().unwrap().to_vec(),
col_idx: csr.indices().to_vec(),
values: csr.data().to_vec(),
}
}
/// Convert to sprs CSR matrix
pub fn to_csr(&self) -> CsMatI<f64, usize> {
CsMatI::new(
(self.n_rows, self.n_cols),
self.row_ptr.clone(),
self.col_idx.clone(),
self.values.clone(),
)
}
/// Number of non-zeros
pub fn nnz(&self) -> usize {
self.values.len()
}
/// Sparsity ratio (nnz / n_elements)
pub fn sparsity(&self) -> f64 {
let total = (self.n_rows * self.n_cols) as f64;
if total > 0.0 {
self.values.len() as f64 / total
} else {
0.0
}
}
}
/// FEM assembler configuration
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AssemblerConfig {
/// Use parallel assembly
pub parallel: bool,
/// Chunk size for parallel assembly
pub chunk_size: usize,
/// Apply reference electrode (average reference)
pub apply_reference: bool,
/// Reference electrode node index (if not average)
pub reference_node: Option<usize>,
}
impl Default for AssemblerConfig {
fn default() -> Self {
Self {
parallel: true,
chunk_size: 1000,
apply_reference: true,
reference_node: None,
}
}
}
/// FEM matrix assembler
#[derive(Debug)]
pub struct FemAssembler {
/// Configuration
config: AssemblerConfig,
/// Mesh reference
mesh: HeadMesh,
/// Conductivity model
conductivity: TissueConductivity,
/// Element stiffness matrices
element_stiffness: Vec<ElementStiffness>,
/// Global stiffness matrix
global_stiffness: Option<CsMatI<f64, usize>>,
}
impl FemAssembler {
/// Create a new assembler
pub fn new(mesh: HeadMesh, conductivity: TissueConductivity, config: AssemblerConfig) -> Self {
Self {
config,
mesh,
conductivity,
element_stiffness: Vec::new(),
global_stiffness: None,
}
}
/// Create with default configuration
pub fn with_defaults(mesh: HeadMesh, conductivity: TissueConductivity) -> Self {
Self::new(mesh, conductivity, AssemblerConfig::default())
}
/// Get mesh reference
pub fn mesh(&self) -> &HeadMesh {
&self.mesh
}
/// Get conductivity model
pub fn conductivity(&self) -> &TissueConductivity {
&self.conductivity
}
/// Compute element stiffness matrices
pub fn compute_element_stiffness(&mut self) -> FemResult<()> {
// Ensure element conductivity tensors are computed
let mut conductivity = self.conductivity.clone();
if conductivity.element_tensors().is_none() {
conductivity.compute_element_tensors(&self.mesh)?;
}
let elements = &self.mesh.elements;
let _nodes = &self.mesh.nodes;
if self.config.parallel {
// Parallel computation
let element_stiffness: Vec<ElementStiffness> = elements
.par_iter()
.enumerate()
.map(|(idx, elem)| {
let tensor = conductivity
.element_tensor(idx)
.cloned()
.unwrap_or_else(|| ConductivityTensor::isotropic(0.33));
ElementStiffness::compute(elem, &self.mesh, &tensor)
})
.collect();
self.element_stiffness = element_stiffness;
} else {
// Sequential computation
self.element_stiffness.clear();
self.element_stiffness.reserve(elements.len());
for (idx, elem) in elements.iter().enumerate() {
let tensor = conductivity
.element_tensor(idx)
.cloned()
.unwrap_or_else(|| ConductivityTensor::isotropic(0.33));
self.element_stiffness
.push(ElementStiffness::compute(elem, &self.mesh, &tensor));
}
}
self.conductivity = conductivity;
Ok(())
}
/// Assemble global stiffness matrix
pub fn assemble_global(&mut self) -> FemResult<()> {
if self.element_stiffness.is_empty() {
self.compute_element_stiffness()?;
}
let n_nodes = self.mesh.n_nodes();
// Use triplet format for assembly
let mut triplets = TriMat::new((n_nodes, n_nodes));
// Estimate capacity (4*4 entries per element)
let estimated_nnz = self.element_stiffness.len() * 16;
triplets.reserve(estimated_nnz);
// Assemble element contributions
for elem_k in &self.element_stiffness {
for i in 0..4 {
let global_i = elem_k.node_indices[i];
for j in 0..4 {
let global_j = elem_k.node_indices[j];
let value = elem_k.matrix[(i, j)];
if value.abs() > 1e-15 {
triplets.add_triplet(global_i, global_j, value);
}
}
}
}
// Convert to CSR
let mut csr: CsMatI<f64, usize> = triplets.to_csr();
// Apply reference electrode constraint
if self.config.apply_reference {
csr = self.apply_reference_constraint(csr)?;
}
self.global_stiffness = Some(csr);
Ok(())
}
/// Apply reference electrode constraint to make system solvable
fn apply_reference_constraint(
&self,
mut stiffness: CsMatI<f64, usize>,
) -> FemResult<CsMatI<f64, usize>> {
let n = stiffness.rows();
if let Some(ref_node) = self.config.reference_node {
// Fix potential at reference node
// Set row and column to zero, diagonal to 1
if ref_node >= n {
return Err(FemError::AssemblyError(
"Reference node index out of bounds".into(),
));
}
// Convert to dense for modification (inefficient but safe)
let mut dense = self.sparse_to_dense(&stiffness);
// Zero out row and column
for j in 0..n {
dense[[ref_node, j]] = 0.0;
dense[[j, ref_node]] = 0.0;
}
dense[[ref_node, ref_node]] = 1.0;
// Convert back to sparse
stiffness = self.dense_to_sparse(&dense);
} else {
// Average reference: deflate matrix
// K_deflated = K - K*v*v^T - v*v^T*K + v^T*K*v * v*v^T
// where v = [1/√n, 1/√n, ..., 1/√n]^T
// This is expensive, so we just add a small regularization
// to the diagonal and use iterative solver with null space handling
// Simple approach: add regularization
let reg = 1e-10;
let mut dense = self.sparse_to_dense(&stiffness);
for i in 0..n {
dense[[i, i]] += reg;
}
stiffness = self.dense_to_sparse(&dense);
}
Ok(stiffness)
}
/// Convert sparse to dense matrix
fn sparse_to_dense(&self, sparse: &CsMatI<f64, usize>) -> Array2<f64> {
let (rows, cols) = (sparse.rows(), sparse.cols());
let mut dense = Array2::zeros((rows, cols));
for (&val, (row, col)) in sparse {
dense[[row, col]] = val;
}
dense
}
/// Convert dense to sparse matrix
fn dense_to_sparse(&self, dense: &Array2<f64>) -> CsMatI<f64, usize> {
let (rows, cols) = dense.dim();
let mut triplets = TriMat::new((rows, cols));
for i in 0..rows {
for j in 0..cols {
let val = dense[[i, j]];
if val.abs() > 1e-15 {
triplets.add_triplet(i, j, val);
}
}
}
triplets.to_csr()
}
/// Get global stiffness matrix
pub fn stiffness_matrix(&self) -> Option<&CsMatI<f64, usize>> {
self.global_stiffness.as_ref()
}
/// Get stiffness matrix in serializable format
pub fn stiffness_matrix_dto(&self) -> Option<StiffnessMatrix> {
self.global_stiffness
.as_ref()
.map(StiffnessMatrix::from_csr)
}
/// Get number of nodes
pub fn n_nodes(&self) -> usize {
self.mesh.n_nodes()
}
/// Get number of elements
pub fn n_elements(&self) -> usize {
self.mesh.n_elements()
}
/// Get assembly statistics
pub fn stats(&self) -> AssemblyStats {
let n_nodes = self.mesh.n_nodes();
let n_elements = self.mesh.n_elements();
let (nnz, sparsity) = if let Some(ref k) = self.global_stiffness {
(k.nnz(), k.nnz() as f64 / (n_nodes * n_nodes) as f64)
} else {
(0, 0.0)
};
AssemblyStats {
n_nodes,
n_elements,
n_element_stiffness: self.element_stiffness.len(),
nnz,
sparsity,
parallel: self.config.parallel,
}
}
}
/// Assembly statistics
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AssemblyStats {
/// Number of mesh nodes
pub n_nodes: usize,
/// Number of elements
pub n_elements: usize,
/// Number of computed element stiffness matrices
pub n_element_stiffness: usize,
/// Number of non-zeros in global matrix
pub nnz: usize,
/// Sparsity ratio
pub sparsity: f64,
/// Whether parallel assembly was used
pub parallel: bool,
}
/// Compute transfer matrix for sensors
///
/// Given electrode positions, compute the interpolation matrix that
/// maps nodal potentials to electrode potentials.
#[derive(Debug, Clone)]
pub struct TransferMatrix {
/// Electrode positions
pub electrode_positions: Vec<Vector3<f64>>,
/// Transfer matrix (n_electrodes x n_nodes)
pub matrix: Array2<f64>,
/// Element indices containing each electrode
pub electrode_elements: Vec<Option<usize>>,
}
impl TransferMatrix {
/// Compute transfer matrix for electrode positions
pub fn compute(mesh: &HeadMesh, electrodes: &[Vector3<f64>]) -> FemResult<Self> {
let n_electrodes = electrodes.len();
let n_nodes = mesh.n_nodes();
let mut matrix = Array2::zeros((n_electrodes, n_nodes));
let mut electrode_elements = Vec::with_capacity(n_electrodes);
for (i, pos) in electrodes.iter().enumerate() {
// Find element containing electrode
if let Some(elem_idx) = mesh.find_element(pos) {
let elem = &mesh.elements[elem_idx];
let bary = compute_barycentric(mesh, elem, pos);
// Interpolate using shape functions
for (j, &node_idx) in elem.nodes.iter().enumerate() {
matrix[[i, node_idx]] = bary[j];
}
electrode_elements.push(Some(elem_idx));
} else {
// Electrode outside mesh - find nearest node
let mut min_dist = f64::MAX;
let mut nearest_node = 0;
for (j, node) in mesh.nodes.iter().enumerate() {
let dist = (node.position - pos).norm();
if dist < min_dist {
min_dist = dist;
nearest_node = j;
}
}
matrix[[i, nearest_node]] = 1.0;
electrode_elements.push(None);
}
}
Ok(Self {
electrode_positions: electrodes.to_vec(),
matrix,
electrode_elements,
})
}
/// Apply transfer matrix to nodal potentials
pub fn apply(&self, nodal_potentials: &ndarray::Array1<f64>) -> ndarray::Array1<f64> {
self.matrix.dot(nodal_potentials)
}
/// Number of electrodes
pub fn n_electrodes(&self) -> usize {
self.electrode_positions.len()
}
}
/// Compute barycentric coordinates in a tetrahedron
fn compute_barycentric(mesh: &HeadMesh, elem: &TetElement, point: &Vector3<f64>) -> [f64; 4] {
let p0 = &mesh.nodes[elem.nodes[0]].position;
let p1 = &mesh.nodes[elem.nodes[1]].position;
let p2 = &mesh.nodes[elem.nodes[2]].position;
let p3 = &mesh.nodes[elem.nodes[3]].position;
// Compute volumes of sub-tetrahedra
let v_total = tetrahedron_volume(p0, p1, p2, p3);
if v_total.abs() < 1e-15 {
return [0.25, 0.25, 0.25, 0.25];
}
let l0 = tetrahedron_volume(point, p1, p2, p3) / v_total;
let l1 = tetrahedron_volume(p0, point, p2, p3) / v_total;
let l2 = tetrahedron_volume(p0, p1, point, p3) / v_total;
let l3 = 1.0 - l0 - l1 - l2;
[l0, l1, l2, l3]
}
/// Compute signed volume of tetrahedron
fn tetrahedron_volume(
p0: &Vector3<f64>,
p1: &Vector3<f64>,
p2: &Vector3<f64>,
p3: &Vector3<f64>,
) -> f64 {
let v1 = p1 - p0;
let v2 = p2 - p0;
let v3 = p3 - p0;
v1.dot(&v2.cross(&v3)) / 6.0
}
#[cfg(test)]
mod tests {
use super::*;
use crate::mesh::TissueLayer;
#[test]
fn test_element_stiffness() {
// Create simple mesh with one element
let mut mesh = HeadMesh::new();
mesh.nodes = vec![
crate::mesh::MeshNode::new(0, 0.0, 0.0, 0.0),
crate::mesh::MeshNode::new(1, 1.0, 0.0, 0.0),
crate::mesh::MeshNode::new(2, 0.0, 1.0, 0.0),
crate::mesh::MeshNode::new(3, 0.0, 0.0, 1.0),
];
let mut elem = crate::mesh::TetElement::new(0, [0, 1, 2, 3], TissueLayer::GrayMatter);
elem.compute_volume(&mesh.nodes);
mesh.elements.push(elem);
let sigma = ConductivityTensor::isotropic(0.33);
let ke = ElementStiffness::compute(&mesh.elements[0], &mesh, &sigma);
// Stiffness matrix should be symmetric
for i in 0..4 {
for j in 0..4 {
assert!((ke.matrix[(i, j)] - ke.matrix[(j, i)]).abs() < 1e-10);
}
}
// Row sums should be zero (conservation)
for i in 0..4 {
let row_sum: f64 = (0..4).map(|j| ke.matrix[(i, j)]).sum();
assert!(row_sum.abs() < 1e-10, "Row {} sum = {}", i, row_sum);
}
}
#[test]
fn test_assembler_creation() {
let mesh = HeadMesh::three_layer_sphere(0.08, 0.007, 0.006, 2, 1).unwrap();
let conductivity = TissueConductivity::default_isotropic();
let assembler = FemAssembler::with_defaults(mesh, conductivity);
assert!(assembler.n_nodes() > 0);
assert!(assembler.n_elements() > 0);
}
#[test]
fn test_global_assembly() {
let mesh = HeadMesh::three_layer_sphere(0.08, 0.007, 0.006, 2, 1).unwrap();
let conductivity = TissueConductivity::default_isotropic();
let mut assembler = FemAssembler::with_defaults(mesh, conductivity);
assembler.assemble_global().unwrap();
let k = assembler.stiffness_matrix().unwrap();
assert!(k.rows() == assembler.n_nodes());
assert!(k.cols() == assembler.n_nodes());
assert!(k.nnz() > 0);
// Check sparsity
let stats = assembler.stats();
assert!(stats.sparsity < 0.1); // Should be very sparse
}
#[test]
fn test_stiffness_symmetry() {
let mesh = HeadMesh::three_layer_sphere(0.08, 0.007, 0.006, 2, 1).unwrap();
let conductivity = TissueConductivity::default_isotropic();
let mut assembler = FemAssembler::with_defaults(mesh, conductivity);
assembler.assemble_global().unwrap();
let k = assembler.stiffness_matrix().unwrap();
// Check symmetry for a few elements
for (&val, (i, j)) in k.iter().take(100) {
// Find K[j,i]
let mut found = false;
for (&v, (ii, jj)) in k.iter() {
if ii == j && jj == i {
assert!(
(val - v).abs() < 1e-10,
"K[{},{}] = {} != K[{},{}] = {}",
i,
j,
val,
j,
i,
v
);
found = true;
break;
}
}
if !found && val.abs() > 1e-15 {
panic!("No symmetric entry found for K[{},{}]", i, j);
}
}
}
#[test]
fn test_transfer_matrix() {
let mesh = HeadMesh::three_layer_sphere(0.08, 0.007, 0.006, 2, 1).unwrap();
// Electrodes on scalp surface
let electrodes = vec![
Vector3::new(0.0, 0.0, 0.093),
Vector3::new(0.093, 0.0, 0.0),
Vector3::new(0.0, 0.093, 0.0),
];
let transfer = TransferMatrix::compute(&mesh, &electrodes).unwrap();
assert_eq!(transfer.n_electrodes(), 3);
assert_eq!(transfer.matrix.nrows(), 3);
assert_eq!(transfer.matrix.ncols(), mesh.n_nodes());
// Each row should sum to 1 (interpolation weights)
for i in 0..3 {
let row_sum: f64 = transfer.matrix.row(i).sum();
assert!((row_sum - 1.0).abs() < 1e-10, "Row {} sum = {}", i, row_sum);
}
}
#[test]
fn test_barycentric_coordinates() {
let mesh = HeadMesh::three_layer_sphere(0.08, 0.007, 0.006, 2, 1).unwrap();
// Find an element and test barycentric at centroid
if let Some(elem) = mesh.elements.first() {
let centroid = elem.centroid(&mesh.nodes);
let bary = compute_barycentric(&mesh, elem, &centroid);
// At centroid, barycentric coords should be approximately equal
let avg = 0.25;
for (i, &b) in bary.iter().enumerate() {
assert!(
(b - avg).abs() < 0.1,
"Barycentric {} = {} (expected ~{})",
i,
b,
avg
);
}
// Should sum to 1
let sum: f64 = bary.iter().sum();
assert!((sum - 1.0).abs() < 1e-10);
}
}
#[test]
fn test_assembly_stats() {
let mesh = HeadMesh::three_layer_sphere(0.08, 0.007, 0.006, 2, 1).unwrap();
let conductivity = TissueConductivity::default_isotropic();
let mut assembler = FemAssembler::with_defaults(mesh, conductivity);
assembler.assemble_global().unwrap();
let stats = assembler.stats();
assert!(stats.n_nodes > 0);
assert!(stats.n_elements > 0);
assert!(stats.nnz > 0);
assert!(stats.sparsity > 0.0);
assert!(stats.sparsity < 1.0);
}
}