// Copyright (c) 2024 RustyTorch++ Team // Licensed under the Apache License, Version 2.0 //! Mesh refinement operations. use crate::error::{FeaResult, MeshError}; use crate::mesh::{Element, ElementType, Node, NodeId}; use std::collections::HashMap; /// Mesh refinement operations. pub struct MeshRefinement; impl MeshRefinement { /// Uniform refinement by subdividing all elements. pub fn uniform_refinement(mesh: &crate::mesh::Mesh) -> FeaResult { let mut refined_mesh = crate::mesh::Mesh::new(mesh.spatial_dimension)?; // Copy existing nodes first let mut node_mapping = HashMap::new(); for (node_id, node) in &mesh.nodes { let new_node = Node { coordinates: node.coordinates.clone(), dofs: node.dofs.clone(), label: node.label.clone(), }; let new_id = refined_mesh.add_node(new_node); node_mapping.insert(*node_id, new_id); } // Track edge midpoints to avoid duplication let mut edge_midpoints: HashMap<(NodeId, NodeId), NodeId> = HashMap::new(); // Helper to get or create edge midpoint let get_edge_midpoint = |n1: NodeId, n2: NodeId, edge_midpoints: &mut HashMap<(NodeId, NodeId), NodeId>, refined_mesh: &mut crate::mesh::Mesh, mesh: &crate::mesh::Mesh| -> NodeId { let edge_key = if n1 < n2 { (n1, n2) } else { (n2, n1) }; if let Some(&mid_id) = edge_midpoints.get(&edge_key) { return mid_id; } // Create midpoint node let node1 = &mesh.nodes[&n1]; let node2 = &mesh.nodes[&n2]; let midpoint_coords = (&node1.coordinates + &node2.coordinates) * 0.5; let mid_node = Node { coordinates: midpoint_coords, dofs: node1.dofs.clone(), // Use same DOF types as parent nodes label: None, }; let mid_id = refined_mesh.add_node(mid_node); edge_midpoints.insert(edge_key, mid_id); mid_id }; // Process each element for element in mesh.elements.values() { match element.element_type { ElementType::Tri3 => { // Triangle subdivision into 4 triangles if element.nodes.len() != 3 { return Err(MeshError::InvalidElement(format!( "Triangle element requires 3 nodes, found {}", element.nodes.len() )) .into()); } let n0 = node_mapping[&element.nodes[0]]; let n1 = node_mapping[&element.nodes[1]]; let n2 = node_mapping[&element.nodes[2]]; // Create edge midpoints let m01 = get_edge_midpoint( element.nodes[0], element.nodes[1], &mut edge_midpoints, &mut refined_mesh, mesh, ); let m12 = get_edge_midpoint( element.nodes[1], element.nodes[2], &mut edge_midpoints, &mut refined_mesh, mesh, ); let m20 = get_edge_midpoint( element.nodes[2], element.nodes[0], &mut edge_midpoints, &mut refined_mesh, mesh, ); // Create 4 new triangles refined_mesh.add_element(Element::new( ElementType::Tri3, vec![n0, m01, m20], element.material_id, )?); refined_mesh.add_element(Element::new( ElementType::Tri3, vec![m01, n1, m12], element.material_id, )?); refined_mesh.add_element(Element::new( ElementType::Tri3, vec![m20, m12, n2], element.material_id, )?); refined_mesh.add_element(Element::new( ElementType::Tri3, vec![m01, m12, m20], element.material_id, )?); } ElementType::Quad4 => { // Quad subdivision into 4 quads if element.nodes.len() != 4 { return Err(MeshError::InvalidElement(format!( "Quad element requires 4 nodes, found {}", element.nodes.len() )) .into()); } let n0 = node_mapping[&element.nodes[0]]; let n1 = node_mapping[&element.nodes[1]]; let n2 = node_mapping[&element.nodes[2]]; let n3 = node_mapping[&element.nodes[3]]; // Create edge midpoints let m01 = get_edge_midpoint( element.nodes[0], element.nodes[1], &mut edge_midpoints, &mut refined_mesh, mesh, ); let m12 = get_edge_midpoint( element.nodes[1], element.nodes[2], &mut edge_midpoints, &mut refined_mesh, mesh, ); let m23 = get_edge_midpoint( element.nodes[2], element.nodes[3], &mut edge_midpoints, &mut refined_mesh, mesh, ); let m30 = get_edge_midpoint( element.nodes[3], element.nodes[0], &mut edge_midpoints, &mut refined_mesh, mesh, ); // Create center point let center_coords = (&mesh.nodes[&element.nodes[0]].coordinates + &mesh.nodes[&element.nodes[1]].coordinates + &mesh.nodes[&element.nodes[2]].coordinates + &mesh.nodes[&element.nodes[3]].coordinates) * 0.25; let center_node = Node { coordinates: center_coords, dofs: mesh.nodes[&element.nodes[0]].dofs.clone(), label: None, }; let center = refined_mesh.add_node(center_node); // Create 4 new quads refined_mesh.add_element(Element::new( ElementType::Quad4, vec![n0, m01, center, m30], element.material_id, )?); refined_mesh.add_element(Element::new( ElementType::Quad4, vec![m01, n1, m12, center], element.material_id, )?); refined_mesh.add_element(Element::new( ElementType::Quad4, vec![center, m12, n2, m23], element.material_id, )?); refined_mesh.add_element(Element::new( ElementType::Quad4, vec![m30, center, m23, n3], element.material_id, )?); } ElementType::Tet4 => { // Tetrahedral subdivision into 8 tets (octahedral subdivision) if element.nodes.len() != 4 { return Err(MeshError::InvalidElement(format!( "Tetrahedral element requires 4 nodes, found {}", element.nodes.len() )) .into()); } let n0 = node_mapping[&element.nodes[0]]; let n1 = node_mapping[&element.nodes[1]]; let n2 = node_mapping[&element.nodes[2]]; let n3 = node_mapping[&element.nodes[3]]; // Create all edge midpoints (6 edges for tet) let m01 = get_edge_midpoint( element.nodes[0], element.nodes[1], &mut edge_midpoints, &mut refined_mesh, mesh, ); let m02 = get_edge_midpoint( element.nodes[0], element.nodes[2], &mut edge_midpoints, &mut refined_mesh, mesh, ); let m03 = get_edge_midpoint( element.nodes[0], element.nodes[3], &mut edge_midpoints, &mut refined_mesh, mesh, ); let m12 = get_edge_midpoint( element.nodes[1], element.nodes[2], &mut edge_midpoints, &mut refined_mesh, mesh, ); let m13 = get_edge_midpoint( element.nodes[1], element.nodes[3], &mut edge_midpoints, &mut refined_mesh, mesh, ); let m23 = get_edge_midpoint( element.nodes[2], element.nodes[3], &mut edge_midpoints, &mut refined_mesh, mesh, ); // Create 8 new tets - 4 corner tets and 4 from octahedron in center // Corner tets refined_mesh.add_element(Element::new( ElementType::Tet4, vec![n0, m01, m02, m03], element.material_id, )?); refined_mesh.add_element(Element::new( ElementType::Tet4, vec![m01, n1, m12, m13], element.material_id, )?); refined_mesh.add_element(Element::new( ElementType::Tet4, vec![m02, m12, n2, m23], element.material_id, )?); refined_mesh.add_element(Element::new( ElementType::Tet4, vec![m03, m13, m23, n3], element.material_id, )?); // Octahedron subdivision into 4 tets refined_mesh.add_element(Element::new( ElementType::Tet4, vec![m01, m02, m03, m13], element.material_id, )?); refined_mesh.add_element(Element::new( ElementType::Tet4, vec![m01, m02, m12, m13], element.material_id, )?); refined_mesh.add_element(Element::new( ElementType::Tet4, vec![m02, m03, m13, m23], element.material_id, )?); refined_mesh.add_element(Element::new( ElementType::Tet4, vec![m02, m12, m13, m23], element.material_id, )?); } ElementType::Hex8 => { // Hexahedral subdivision into 8 hexahedra if element.nodes.len() != 8 { return Err(MeshError::InvalidElement(format!( "Hexahedral element requires 8 nodes, found {}", element.nodes.len() )) .into()); } // Map original nodes let original_nodes: Vec<_> = element.nodes.iter().map(|&n| node_mapping[&n]).collect(); // Create all edge midpoints (12 edges for hex) let mut edge_mids = Vec::new(); let edge_pairs = [ (0, 1), (1, 2), (2, 3), (3, 0), // Bottom face (4, 5), (5, 6), (6, 7), (7, 4), // Top face (0, 4), (1, 5), (2, 6), (3, 7), // Vertical edges ]; for &(i, j) in &edge_pairs { let mid = get_edge_midpoint( element.nodes[i], element.nodes[j], &mut edge_midpoints, &mut refined_mesh, mesh, ); edge_mids.push(mid); } // Create face centers (6 faces) let face_indices = [ [0, 1, 2, 3], // Bottom [4, 5, 6, 7], // Top [0, 1, 5, 4], // Front [2, 3, 7, 6], // Back [0, 3, 7, 4], // Left [1, 2, 6, 5], // Right ]; let mut face_centers = Vec::new(); for face in &face_indices { let mut center = nalgebra::DVector::zeros(mesh.spatial_dimension); for &idx in face { center += &mesh.nodes[&element.nodes[idx]].coordinates; } center /= 4.0; let face_node = Node { coordinates: center, dofs: mesh.nodes[&element.nodes[0]].dofs.clone(), label: None, }; face_centers.push(refined_mesh.add_node(face_node)); } // Create volume center let mut vol_center = nalgebra::DVector::zeros(mesh.spatial_dimension); for &node_id in &element.nodes { vol_center += &mesh.nodes[&node_id].coordinates; } vol_center /= 8.0; let vol_center_node = Node { coordinates: vol_center, dofs: mesh.nodes[&element.nodes[0]].dofs.clone(), label: None, }; let vol_center_id = refined_mesh.add_node(vol_center_node); // Create 8 new hexahedra // This is complex - each original vertex gets a new hex let new_hex_nodes = [ // Hex at vertex 0 vec![ original_nodes[0], edge_mids[0], face_centers[0], edge_mids[3], edge_mids[8], face_centers[2], vol_center_id, face_centers[4], ], // Hex at vertex 1 vec![ edge_mids[0], original_nodes[1], edge_mids[1], face_centers[0], face_centers[2], edge_mids[9], face_centers[5], vol_center_id, ], // Hex at vertex 2 vec![ face_centers[0], edge_mids[1], original_nodes[2], edge_mids[2], vol_center_id, face_centers[5], edge_mids[10], face_centers[3], ], // Hex at vertex 3 vec![ edge_mids[3], face_centers[0], edge_mids[2], original_nodes[3], face_centers[4], vol_center_id, face_centers[3], edge_mids[11], ], // Hex at vertex 4 vec![ edge_mids[8], face_centers[2], vol_center_id, face_centers[4], original_nodes[4], edge_mids[4], face_centers[1], edge_mids[7], ], // Hex at vertex 5 vec![ face_centers[2], edge_mids[9], face_centers[5], vol_center_id, edge_mids[4], original_nodes[5], edge_mids[5], face_centers[1], ], // Hex at vertex 6 vec![ vol_center_id, face_centers[5], edge_mids[10], face_centers[3], face_centers[1], edge_mids[5], original_nodes[6], edge_mids[6], ], // Hex at vertex 7 vec![ face_centers[4], vol_center_id, face_centers[3], edge_mids[11], edge_mids[7], face_centers[1], edge_mids[6], original_nodes[7], ], ]; for nodes in &new_hex_nodes { refined_mesh.add_element(Element::new( ElementType::Hex8, nodes.clone(), element.material_id, )?); } } _ => { // For other element types, just copy as-is refined_mesh.add_element(element.clone()); } } } Ok(refined_mesh) } /// Adaptive refinement based on error indicators. pub fn adaptive_refinement( mesh: &crate::mesh::Mesh, error_indicators: &HashMap, threshold: f64, ) -> FeaResult { let mut refined_mesh = crate::mesh::Mesh::new(mesh.spatial_dimension)?; // Copy all nodes first let mut node_mapping = HashMap::new(); for (node_id, node) in &mesh.nodes { let new_node = Node { coordinates: node.coordinates.clone(), dofs: node.dofs.clone(), label: node.label.clone(), }; let new_id = refined_mesh.add_node(new_node); node_mapping.insert(*node_id, new_id); } // Track edge midpoints let mut edge_midpoints: HashMap<(NodeId, NodeId), NodeId> = HashMap::new(); // Process each element based on error indicator for (element_id, element) in &mesh.elements { let error = error_indicators.get(element_id).copied().unwrap_or(0.0); if error > threshold { // Refine this element // For simplicity, use uniform refinement logic per element // In practice, this would be more sophisticated if element.element_type == ElementType::Tri3 { // Triangle refinement (same as uniform case) let n0 = node_mapping[&element.nodes[0]]; let n1 = node_mapping[&element.nodes[1]]; let n2 = node_mapping[&element.nodes[2]]; // Helper to get/create midpoint let mut get_mid = |i: usize, j: usize| -> NodeId { let ni = element.nodes[i]; let nj = element.nodes[j]; let key = if ni < nj { (ni, nj) } else { (nj, ni) }; *edge_midpoints.entry(key).or_insert_with(|| { let mid_coords = (&mesh.nodes[&ni].coordinates + &mesh.nodes[&nj].coordinates) * 0.5; let mid_node = Node { coordinates: mid_coords, dofs: mesh.nodes[&ni].dofs.clone(), label: None, }; refined_mesh.add_node(mid_node) }) }; let m01 = get_mid(0, 1); let m12 = get_mid(1, 2); let m20 = get_mid(2, 0); // Create 4 refined triangles refined_mesh.add_element(Element::new( ElementType::Tri3, vec![n0, m01, m20], element.material_id, )?); refined_mesh.add_element(Element::new( ElementType::Tri3, vec![m01, n1, m12], element.material_id, )?); refined_mesh.add_element(Element::new( ElementType::Tri3, vec![m20, m12, n2], element.material_id, )?); refined_mesh.add_element(Element::new( ElementType::Tri3, vec![m01, m12, m20], element.material_id, )?); } else { // For other element types or below threshold, copy as-is let new_nodes: Vec<_> = element.nodes.iter().map(|&n| node_mapping[&n]).collect(); refined_mesh.add_element(Element::new( element.element_type, new_nodes, element.material_id, )?); } } else { // Copy element without refinement let new_nodes: Vec<_> = element.nodes.iter().map(|&n| node_mapping[&n]).collect(); refined_mesh.add_element(Element::new( element.element_type, new_nodes, element.material_id, )?); } } Ok(refined_mesh) } }