JAX-FEM: A differentiable GPU-accelerated 3D finite element solver for automatic inverse design and mechanistic data science
arXiv:2212.00964 · doi:10.1016/j.cpc.2023.108802
Abstract
This paper introduces JAX-FEM, an open-source differentiable finite element method (FEM) library. Constructed on top of Google JAX, a rising machine learning library focusing on high-performance numerical computing, JAX-FEM is implemented with pure Python while scalable to efficiently solve problems with moderate to large sizes. For example, in a 3D tensile loading problem with 7.7 million degrees of freedom, JAX-FEM with GPU achieves around 10 acceleration compared to a commercial FEM code depending on platform. Beyond efficiently solving forward problems, JAX-FEM employs the automatic differentiation technique so that inverse problems are solved in a fully automatic manner without the need to manually derive sensitivities. Examples of 3D topology optimization of nonlinear materials are shown to achieve optimal compliance. Finally, JAX-FEM is an integrated platform for machine learning-aided computational mechanics. We show an example of data-driven multi-scale computations of a composite material where JAX-FEM provides an all-in-one solution from microscopic data generation and model training to macroscopic FE computations. The source code of the library and these examples are shared with the community to facilitate computational mechanics research.
References in corpus (4)
Cited by in corpus (13)
- XLB: A differentiable massively parallel lattice Boltzmann library in Python
- Tensor-decomposition-based A Priori Surrogate (TAPS) modeling for ultra large-scale simulations
- Gradient-free neural topology optimization: Towards effective fracture-resistant designs
- Scientific Machine Learning of Flow Resistance Using Universal Shallow Water Equations with Differentiable Programming
- An -adaptive finite element method using neural networks for parametric self-adjoint elliptic problem
- COMMET: orders-of-magnitude speed-up in finite element method via batch-vectorized neural constitutive updates
- GeoWarp: An automatically differentiable and GPU-accelerated implicit MPM framework for geomechanics based on NVIDIA Warp
- Lattice Discrete Particle Model (LDPM): Comparison of Various Time Integration Solvers and Implementations
- Active learning with physics-informed neural networks for optimal sensor placement in deep tunneling through transversely isotropic elastic rocks
- Gradient-based optimization of scatterer arrangements based on the T-matrix method
- Material-agnostic temperature field prediction for metal additive manufacturing via a parametric PINN framework
- JAX-AMG: A GPU-Accelerated Differentiable Sparse Linear Solver Library for JAX
- Iterative Learning Control of the Cooling Rate in a Dual-Laser Powder Bed Fusion Process