Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
@@ -0,0 +1,674 @@
//! 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);
}
}