//! 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, /// 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, /// CSR column indices pub col_idx: Vec, /// CSR values pub values: Vec, } impl StiffnessMatrix { /// Create from sprs CSR matrix pub fn from_csr(csr: &CsMatI) -> 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 { 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, } 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, /// Global stiffness matrix global_stiffness: Option>, } 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 = 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 = 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, ) -> FemResult> { 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) -> Array2 { 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) -> CsMatI { 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> { self.global_stiffness.as_ref() } /// Get stiffness matrix in serializable format pub fn stiffness_matrix_dto(&self) -> Option { 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>, /// Transfer matrix (n_electrodes x n_nodes) pub matrix: Array2, /// Element indices containing each electrode pub electrode_elements: Vec>, } impl TransferMatrix { /// Compute transfer matrix for electrode positions pub fn compute(mesh: &HeadMesh, electrodes: &[Vector3]) -> FemResult { 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) -> ndarray::Array1 { 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; 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, p1: &Vector3, p2: &Vector3, p3: &Vector3, ) -> 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, ¢roid); // 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); } }