E3J: An Efficient and Open-Source Backend for Euclidean Equivariant Operations on GPU and TPU
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.