Fast open source backend accelerates equivariant operations on gpu and tpu
E3J: An Efficient and Open-Source Backend for Euclidean Equivariant Operations on GPU and TPU
Machine LearningDistributed, Parallel, and Cluster ComputingMathematical Software
Summary
Many AI applications involve math that doesn’t change when you rotate or move things around, called Euclidean equivariance. The authors built e3j, an open-source software that runs these special math operations very fast on modern processors like GPUs and TPUs. They showed it works better than other tools on tasks like simulating water molecules. Their work helps developers run complex models much faster while making the code easy to use and share.
What this means in practice
- •For machine learning engineers: Run geometric deep learning models with Euclidean symmetry much faster on GPUs and TPUs for physical simulations and molecular modeling.
- •For high performance computing teams: Accelerate scientific computing workloads involving tensor products and equivariant operations using an efficient open-source backend compatible with GPUs and TPUs.
Authors
Olivier Peltre, Armand Picard, Adrien Pichard, Miguel Bragança, Luca Giacomoni, Valentin Heyraud, Zachary Weller-Davies, Christoph Brunken, Jules Tilly
Abstract
We present e3j, a fast Euclid-equivariance backend for geometric deep learning applications with JAX bindings for GPU and TPU. Leveraging both optimized CUDA and Pallas kernels and algorithmic improvements, the library achieves state-of-the-art throughput and runtime on both forward and backward paths. On a machine learning interatomic potential (MLIP) use case, it outperforms established backends, measuring up to 34% speed-up over cuEquivariance on water box NPT simulation using MACE, while remaining fully open source. E3j achieves over 80% efficiency over the H100 maximum memory bandwidth on tensor product operations, and in many cases more than doubles throughput of message passing convolutions forward compared to previously available backends. In addition, with the release of dedicated Pallas TPU kernel, e3j opens the possibility of large scale equivariant deep learning workloads on TPU architectures, which has so far been difficult to achieve. Our benchmarks show that e3j also achieves over 80% of a TPUv6e memory bandwidth, up to one order of magnitude more than e3nn-jax. The library is available on GitHub, PyPI and is released under an open source Apache 2.0 license.