Files
rustytorch/crates/specialized/rtx-fea/src/mesh/topology/topology_refinement.rs
T
2026-03-04 00:08:42 +00:00

614 lines
24 KiB
Rust

// 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<crate::mesh::Mesh> {
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<crate::mesh::ElementId, f64>,
threshold: f64,
) -> FeaResult<crate::mesh::Mesh> {
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)
}
}