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.
Figures & tables
Channels
ℓmax
Implementation
Det.
64
128
256
512
1024
Forward pass
2
e3nn (f32) †
00 89 ± 0 1
00 90 ± 0 1
00 90 ± 0 0
00 92 ± 0 0
00 90 ± 0 0
CuEquivariance
0 779 ± 0 1
0 833 ± 15
0 786 ± 0 4
0 807 ± 11
0 755 ± 0 0
OpenEquivariance
✓
0 871 ± 0 3
1092 ± 0 8
0 760 ± 0 8
0 631 ± 0 5
0 652 ± 0 7
OpenEquivariance
0 348 ± 0 1
0 343 ± 0 1
0 370 ± 0 1
0 354 ± 0 2
0 443 ± 0 1
Table 1: Maximal convolution throughput (GB/s) by channel count on GPU. End-to-end throughput is reported at the maximal power-of-two node count fitting on the NVIDIA H100 HBM3, with 45 average neighbors. This keeps the product of node count with channel count fixed to 222 at ℓmax=2 , and to 221 at ℓmax=3 for fused kernels. The e3nn baseline is kept to illustrate the yield of a pure JAX implementation of the same operation, but message materialization overflows memory at 218 and 217 respectively. The Pallas GPU kernel does not support channel counts below 128 yet (minimal block size imposed by Pallas).
Backend
Deterministic
Open-source
ms/step
ns/day
e3j (CUDA)
✓
✓
8.821 ± 0.045
9.80 ± 0.05
OpenEquivariance
✓
✓
16.137 ± 1.063
5.38 ± 0.35
cuEquivariance
6.586 ± 0.013
13.12 ± 0.03
e3j (Pallas GPU)
✓
4.906 ± 0.014
17.61 ± 0.05
OpenEquivariance
✓
10.041 ± 0.040
8.60 ± 0.03
e3nn
✓
68.298 ± 0.063
1.27 ± 0.00
Table 2: End-to-end NPT performance of a MACE model on a 25Å water box (GPU). Runtimes are reported for 100ps long NPT simulations with a Monte Carlo barostat, Hyperparameters are from the MACE (a) variant of table 3 , notably correlation=2 and node_symmetry=2 . Only the convolution block is dispatched to dedicated kernels matching the reference implementation numerically. The initial structure (solvated 2-methyl-butane, equilibrated with a classical force field) consists of 1503 atoms and 10% initial edge padding (83,325 static edge count).
Appendix figures & tables2 assets
Supplementary material from the paper’s appendix.
Appendix
MACE (a)
num_layers
2
num_channels
128*
correlation
2
node_symmetry
2
l_max
3
cutoff_angstrom
5
Appendix
Table 3: Hyperparameters used in end-to-end MLIP benchmarks.
pass
B
buffers
peak VMEM
fwd
128
4.3
13.1 ( 82% )
bwd
32
5.4
10.1 ( 63% )
Appendix
Table 4: VMEM usage on a TPU v4 TensorCore of the fused TPU kernel at lmax=3 , C=128 , fp32, in MiB against the 16 MiB per-core budget; buffers counts the kernel’s own arrays, peak VMEM adds the compiler’s spill.