Preprint
E3J: A N E FFICIENT AND O PEN -S OURCE BACKEND FOR E UCLIDEAN E QUIVARIANT O PERATIONS ON GPU AND TPU Olivier Peltre InstaDeep
Armand Picard InstaDeep
arXiv:2609.35099v1 [cs.LG] 28 Sep 2026
Valentin Heyraud InstaDeep
Adrien Pichard InstaDeep
Zachary Weller-Davies InstaDeep
Miguel Bragança InstaDeep
Luca Giacomoni∗ Prima Mente
Christoph Brunken InstaDeep
Jules Tilly InstaDeep
A BSTRACT 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.
1
I NTRODUCTION
Euclid-equivariant operations are a core building block for many geometric deep learning applications. By providing exact equivariance guarantees, they have been shown to allow aggregation of high order geometric features in low data regime and without need for data augmentation [1, 2]. Their use has been explored across many fields that rely on geometric data, and in particular physical sciences, such as molecular dynamics [2–5], computational fluid dynamics (CFD) [6, 7], or electronic wavefunction and densities [8]. In this work, we present e3j, an open-source Euclid-equivariance backend providing all the core building blocks for E(3)-equivariant neural networks. The library is designed to facilitate constructions of these networks through a modular interface, but crucially is used as a means to integrate dedicated accelerated kernels into model architectures. To that end, the library includes kernels written both in Pallas and in CUDA, with the aim of deploying workloads across multiple types of architectures, including TPUs and GPUs. Namely, we present three sets of kernels backed by novel algorithms: • CUDA kernels: Fully deterministic equivariant operations for GPU, focused on broad compatibility across JAX versions through ahead-of-time (AOT) compilation. • Pallas GPU kernels: Focused on performance, setting state-of-the-art throughput at time of writing in most regimes tested through just-in-time (JIT) compilation by the Mosaic compiler. Sacrifices determinism on the backward pass, compatibility with some older JAX versions, and support for some operations (e.g. low number of channels). ∗
Contributed while at InstaDeep.
1
Preprint
• Pallas TPU kernels: Focused on delivering optimized throughput for TPUs. We include and benchmark kernels specifically for both the equivariant tensor product, and for the full message passing convolution. In order to illustrate the capabilities of the library in realistic workloads, we focus on applications in the field of machine learning interatomic potentials (MLIPs) which are likely among the most relevant use case examples for the technology as (1) they require a complete absence of systematic bias from frame orientation to produce stable, long time molecular simulations, and (2) unlike other fields of application such as CFD, they rely on nearly perfectly equivariant training data. The library achieves throughputs comparable to those of closed source alternatives, for example reaching 2.7 TB/s on H100 GPU (81% of memory bandwidth) on the forward tensor product operation, and 1.33 TB/s on TPUv6.2e (also 81% of memory bandwidth). In order to illustrate the practical benefits of e3j, we also benchmarked end-to-end integration within the MACE [3, 4] and NequIP [2] models, notably providing about 5x force inference speedup over e3nn on TPU. On GPU, e3j provides 2x speedup in simulations over the best open-source backend OpenEquivariance, and 34% speedup over the proprietary cuEquivariance backend, reaching up to 64 thousands atoms (with ∼ 40 connectivity degree) on MACE before overflowing the H100 memory. Details of these results are presented in Section 3. For theoretical details about the operations included in the library, readers can refer to Appendix A.3. The details of our novel algorithmic and methodological developments are presented in Appendix E.
2
R ELATED W ORK
Engineering work: There has been considerable effort in building parameterized equivariant transforms accessible to the machine learning (ML) and scientific computing ecosystem. Generally, this consists in extending standard ML frameworks such as JAX [9] and PyTorch [10] to benefit from their automatic differentiation (AD) support, so that equivariant building blocks can be seamlessly integrated within larger workflows on accelerated hardware such as GPUs and TPUs. These extensions may be defined either within the AD framework itself (JAX / PyTorch), or via the lower-level definition of ad-hoc primitives in the CUDA language1 [11]. The main examples of relevant libraries in this space are: • e3nn [12, 13], one of the first and most feature-rich Euclid equivariance libraries, consisting of two sibling packages written in Torch and JAX, respectively. It has been used to construct widely used MLIP architectures such as MACE [3, 4] and NequIP [2], which can learn to predict energies and forces at quantum-levels of accuracy, thus disrupting the speed-accuracy compromise of MD simulations [14]. • e3x [15], an open source Euclid equivariance package written in JAX [9] which is slightly less flexible than e3nn due to its stricter data model, but offers efficient bilinear projections with cubic L scaling [16]. • OpenEquivariance [17], an open source CUDA kernel generation package with Torch and JAX bindings. It delivers state-of-the-art performance on focused operations (tensor product, message-passing convolution) with dedicated double-backward kernels to optimize training and Hessian inference. It is meant as a drop-in replacement for specific e3nn operations on which it provides speedups of an order of magnitude. • cuEquivariance (cuEq), a proprietary NVIDIA® package with a toolset already as complete as e3nn, providing efficient CUDA kernels for tensor products and polynomial evaluation along with JAX and Torch bindings. cuEquivariance is the most efficient backend, with an open-source API that however depends on closed-source kernels. The above list includes the most feature-complete packages viable for the construction of larger equivariant neural networks such as MLIPs. Efficient, near-optimal (open-source or closed-source) solutions therefore exist to compute Clebsch-Gordan tensor products on NVIDIA GPUs. What best distinguishes e3j from recent work is that e3j provides a standalone and platform-agnostic API in JAX for GPU and TPU execution. 1 While the CUDA language from NVIDIA® extends C++ to let users write general purpose programs on their GPUs, there is no low-level public kernel language for TPUs.
2
Preprint
Theoretical work: Several alternative approaches to improve equivariant operations in geometric deep learning have come from the fundamental side. Naively a Clebsch–Gordan tensor product scales as O(L6 ), or as O(L5 ) if one exploits sparsity. Mathematical work has shown that cubic scaling with L can be obtained for equivariant tensor products, though it is worth noting that these always come with compromise in terms of applicability or expressivity [18]. Some methods demonstrate cubic scaling with L on arbitrary inputs. Gaunt tensor product formulas [19] are morally similar to performing element-wise products instead of discrete convolution via reciprocal Fourier transforms on the input and outputs. While the initial formulation of the Gaunt tensor product (GTP) however fails to fully reproduce the Clebsch-Gordan tensor product, as it does not incorporate anti-symmetric elements, limitations were remediated by a series of papers which formulate and incorporate an anti-symmetric counterpart called Vector Signal Tensor Product (VSTP) [18, 20], and later [21, 22] providing a closed form formula for efficient complete simulation of Clebsch-Gordan tensor products. Let us also mention the matrix tensor product formulation of [16] which also has O(L3 ) scaling and is implemented in e3x, providing order of magnitude speedups at L = 10 over the equivalent e3nn implementation. All of these methods however do result in automatic collapse of output multiplicity, a more efficient implementation which however does lead to a loss in expressivity [18]. Other methods specialize in delivering efficiency gains in the special case of a tensor product between an arbitrary feature vector of irreducible representations and equivariant features obtained by harmonic embeddings of an input vector. This particular case remains dominant in the literature, constituting one of the core building blocks of the original Tensor Field Network [1], later used in NequIP and MACE. A first example is the SO(2) convolution, presented in the equivariant Spherical Channel Network (eSCN) [23] and notably used in eSEN [24] and UMA [5], which defines a frame a reference from each edge and rotates the corresponding tensor product operands using the Wigner-D matrices. This results in harmonic features collapsing to m = 0 across all L, removing the need to sum over m indices of the harmonic features when performing the tensor product. This, combined with the additional use of symmetries in the Clebsch-Gordan tensor product achieves a cubic scaling in L. A second example was presented in the E2former model [25] (later combined with SO(2) convolutions by the same team [26]) where the projected displacement vectors are replaced with the difference of projected input positions into harmonic features. Using a Binomial expansion, the authors show that one can construct a messaging passing block solely relying on node-wise tensor products. While message passing scaling remains unchanged, the number of tensor products (expected to be among the most costly operations) now scales with the number of nodes rather than the number of edges. It is also worth noting that while theoretical work often focuses on tensor product scaling in L, most MLIP applications are in practice interested in the scaling in N but at fixed L (usually 2 or 3)2 . Improving the scaling with L typically requires clever re-parameterizations of equivariant features, and incurs a practical overhead that may not prove beneficial over an efficient implementation of the full Clebsch-Gordan tensor product in the small L regime. We have conducted benchmarks of our backends against methods based on Gaunt tensor products and its extensions, e3x methods and on SO(2) convolution. These are presented in Sec. 3 and in more detail in Appendices B.1, B.2 and B.3 respectively.
3
R ESULTS
The e3j package consists of a Python API targeting the JAX backend, alongside CUDA and Pallas kernel implementations for performance critical operations (tensor product, message-passing). The JAX framework [9] enables seamless integration within larger programs that can be just-in-time (JIT) compiled to XLA (for Accelerated Linear Algebra), as are typically all recent MLIP model implementations. In addition to atomic, module-wise benchmarks against reference E3 backends, we used the open-source mlip library [14] as reference and starting point to estimate the end-to-end speedups e3j integration may provide in the MACE and NequIP equivariant architectures [2, 3]. The details of the implementation method, algorithms, and differentiation rules can be found in Appendix E. 2
Applications such as signal processing or meteorology are in contrast interested in high maximal degree L.
3
Preprint
It is worth noting that our JAX primitives bound to custom kernels are all made infinitely differentiable via recursive AD rules, and compatible with other higher-order JAX transforms (vmap, shard_map,...) to provide SPMD execution on GPU and TPU, an essential feature for MLIP training workflows on a large data scale. Architecture requirements for GPU and TPU however largely differ beyond that point, necessarily resulting in different algorithms. We have added a commentary on the GPU / TPU differences in appendix F. In order to assess the performance of the library, we focus on two types of benchmarking (further discussion on the benchmark details can be found in Appendix D): • Module specific: We benchmark the performance critical components of any Euclidean equivariant library, namely: – Tensor product: Bilinear coupling of latent equivariant features with static ClebschGordan coefficients. In general these are performed per edge and per channel. – Message passing: Aggregates the edge-wise features, usually computed through combination of harmonics projection, tensor product and linear or scalar mixing. The message passing aggregation tends to be the operational bottleneck once efficient tensor product operations are implemented. • End-to-end: To determine the overall relevance of the library in a complete workflow compared to alternatives, we also test two popular MLIP architectures, MACE [3] and NequIP [2]. To maintain comparability, we connected the full benchmark in a fork of the open-source mlip library [14], making the choice of backend the only variable differing in each run. Note that these benchmarks are for illustration only and that applicability of e3j is not restricted to these two architectures (nor to MLIPs in general), a faithful and exhaustive comparison across all possible models and fields of application is not in scope of this work, for obvious reasons. Experiments were conducted running JAX (v0.11.1) on an NVIDIA (R) H100-HBM3 GPU with 3.35 TB/s HBM, and on a TPUv6e Trillium with 1.64 TB/s HBM. We benchmark seven different backends, five for GPU, and two for TPU. On GPU, we compare the following backends with our CUDA and Pallas GPU kernels: • cuEquivariance (v0.11.1): Proprietary tensor product and message-passing convolution kernels of NVIDIA. • OpenEquivariance (v0.7.0): We use JAX bindings to their equivalent tensor product and message-passing convolution kernels, which are JIT compiled from open-source CUDA kernels with NVRTC [17]. Note that OpenEquivariance provides two distinct convolution binaries, a deterministic one for graphs with edges sorted by receiver node index, and a non-deterministic one. • e3nn-jax (v0.21.1): The reference e3nn-jax library sets the performance threshold obtained by a pure JAX implementation, JIT compiled by XLA, but without any specific low level optimization. It should therefore only be considered as an illustrative reference. The e3nn label uses the default half-precision for matmul, while e3nn_f32 enforces single-precision in unit benchmarks. On TPU, we compare our Pallas TPU kernels with XLA compiled e3nn-jax equivalent implementations of the full tensor product and message-passing convolution. Additional benchmark results can be found in appendix C, while comparisons with JAX implementations of cubically scaling algorithms (SO2, Gaunt/VSTP, e3x) are grouped in appendix B. 3.1
U NIT BENCHMARKS
This section provides efficiency comparisons of the currently implemented e3j kernels with available baselines in typical regimes. The tensor product operation is a full Clebsch-Gordan tensor product, satisfying the so-called universal property3 of tensor products. The message-passing convolution op3
Any bilinear map b : (X, Y ) → Z factors through the tensor product space X ⊗ Y as a linear map b̃ : X ⊗ Y → Z.
4
Preprint
eration accumulates edge-wise tensor products on receiver nodes, and consists of a typical bottleneck in MLIP architectures [23, 24, 27]. Runtimes were measured on XLA compiled, numerically equivalent implementations, averaging the fastest 20% of 100 runs, and disabling the Python garbage collector using the timeit module. The backward pass consists of the isolated, XLA compiled vector-jacobian product (VJP). Our metric of interest is throughput, i.e. the total amount of input/output bytes processed in the operation per unit of time. The peak throughput can be directly compared with the GPU / TPU maximum bandwidth to provide a meaningful speed relative to the device, independent of I/O size. In the backward pass benchmarks, only input primals, output cotangents and input cotangents were considered in the VJP accounting, excluding any saved residuals. 3.1.1
T ENSOR P RODUCT
We developed dedicated tensor product kernels using CUDA, Pallas GPU and Pallas TPU, with results presented in Fig. 1. Our results show that the Clebsch-Gordan tensor product (CGTP) operation can reach up to 80% HBM in the forward and backward passes on both GPU and TPU. This means that CGTP in the small L ≤ 3 regime is not a compute bottleneck in itself, while slightly larger degrees L = 4, 5 . . . remain amenable to further engineering optimizations that were not prioritized at this time. On GPU, our Pallas GPU kernel performs almost exactly on par with cuEquivariance, while our CUDA kernel performs similarly on the forward pass and slightly below Pallas GPU / cuEquivariance on the backward pass. All three outpace OpenEquivariance. On TPU our Pallas kernel achieves one order of magnitude higher throughput than the previously available e3nn_jax. TensorProduct (ℓmax = 3, mul = 256) Forward
Backward
Throughput GB/s
TPU v6e
103
Library e3j (CUDA) e3j (Pallas TPU) 10
e3j (Pallas GPU)
2
e3nn e3nn (f32) cuequivariance openequivariance
Throughput GB/s
GPU H100
Reference TPU v6e HBM
103
GPU H100 HBM
102
102
104
103
105
102
Batch size
104
103 Batch size
Figure 1: Throughputs (GB/s) for the tensor product operation on TPU (top) and GPU (bottom). The e3j CUDA implementation is compared against cuEq (jax), and e3nn_jax (single precision and default half precision). The e3j Pallas implementation is compared against the only available baseline e3nn_jax for the universal ("full") Clebsch-Gordan tensor product. Results for ℓmax = 2 and ℓmax = 3 and multiplicities up to 256 channels, alongside the (a) theoretical maximum throughput of the NVIDIA H100 GPU (∼ 3.35 TB/s) and (b) theoretical maximum throughput of a TPU v6e (∼ 1.64 TB/s) on which the benchmarks were performed. When lines are not complete, it indicates that the library reached the memory bound. Note that our tensor product primitive is exposed as a generic bilinear coupling of operands with an arbitrary sparse COO array of coefficients. As presented in appendix B.1, this allows us to further benchmark our kernels against Gaunt [19] and Vector Signal Tensor products [18, 20–22] on 5
Preprint
symmetric and skew-symmetric paths respectively. We find (see figure 5) that for L ≤ 3, our kernels are 2 to 4 times faster than the Gaunt tensor product on symmetric paths, with crossing at L = 6 for the forward and L = 5 for the backward. For the skew-symmetric paths, our kernels outperform the VSTP by a factor of 4 to 6 for L ≤ 3, and the crossing occurs one degree later due to the extra cost of the vector-valued operations in the VSTP. While we believe our implementations of GTP and VSTP to be efficient, dedicated kernel optimization of these operations could improve the relative results. 3.1.2
M ESSAGE PASSING C ONVOLUTION
We present results for our three sets of Message Passing Convolution kernels: CUDA, Pallas GPU and Pallas TPU. On GPU, the Pallas GPU dominates the benchmarks, with over two times the throughput of cuEquivariance in the forward pass under most settings, and a slightly higher throughput in the backward pass. It is worth noting that our Pallas GPU kernel is not (yet) compatible with channel counts lower than 128. The CUDA kernels are broadly on par with OpenEquivariance, both having significantly lower throughput than Pallas GPU and cuEquivariance. These results are presented in Tab. 1. On TPU, our convolution kernels provide well over one order of magnitude speedups on the forward, and nearly one order of magnitude on the backward over e3nn-jax. An overview of the results is presented in Fig. 2. Our convolution kernels further provide the ability to skip padding edges, which typically all point to the same padding node4 in most JAX-based graph neural network frameworks [28]. This feature avoids the potentially significant overhead of aggregating messages on a padding node of unusually high valency, as illustrated by the simulation runtimes summarized in table 2. Implementation details are found in Appendix E. MessagePassing (ℓmax = 3, mul = 256) Forward
Backward
Throughput GB/s
TPU v6e
103
Library e3j (CUDA)
102
e3j (Pallas TPU) e3j (Pallas GPU) † e3nn e3nn (f32) cuequivariance †
101
openequivariance openequivariance † †
non-deterministic
Throughput GB/s
GPU H100
103 Reference TPU v6e HBM GPU H100 HBM 102
101
102
103 Number of nodes
104
102
103 Number of nodes
104
Figure 2: Throughputs (GB/s) for the convolution operation on TPU (top) and GPU (bottom). Here we present results for ℓmax = 3 256 channels, alongside the theoretical maximum throughput of a TPU v6e Trillium (1.64 TB/s) and the theoretical maximum throughput of the NVIDIA H100 HBM3 GPU (∼ 3.35 TB/s) on which the benchmarks were performed. When lines are not complete, it indicates that the library reached the memory bound of the device.
4 Given the graph topology is dynamic, while XLA compiled programs expect static shapes, simulations typically add a single padding node to the atom list while managing a fixed-size buffer of edges, only expanded when the real connecting edges overflow the neighbor list (a rare event triggering re-compilation).
6
Preprint
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). Channels ℓmax
Implementation
Det.
64
128
256
512
1024
89 ± 1 779 ± 1 871 ± 3 348 ± 1 625 ± 2 –
90 ± 1 833 ±15 1092 ± 8 343 ± 1 955 ± 4 1621 ± 9
90 ± 0 786 ± 4 760 ± 8 370 ± 1 913 ± 4 1669 ±12
92 ± 0 807 ±11 631 ± 5 354 ± 2 956 ± 2 1716 ±29
90 ± 0 755 ± 0 652 ± 7 443 ± 1 880 ± 5 1871 ±41
73 ± 0 829 ±11 856 ± 2 265 ± 1 389 ± 0 –
72 ± 1 1072 ± 3 709 ± 1 268 ± 1 569 ± 2 1818 ± 6
75 ± 0 662 ± 1 612 ± 1 303 ± 1 573 ± 1 1859 ±14
75 ± 1 632 ± 1 565 ± 2 323 ± 0 562 ± 2 2074 ±14
73 ± 1 629 ± 0 347 ± 5 303 ± 4 538 ± 2 2067 ±19
Forward pass
2
3
e3nn (f32)† CuEquivariance OpenEquivariance OpenEquivariance e3j (CUDA) e3j (Pallas GPU) e3nn (f32)† CuEquivariance OpenEquivariance OpenEquivariance e3j (CUDA) e3j (Pallas GPU)
✓ ✓ ✓
✓ ✓ ✓
Backward pass †
2
3
e3nn (f32) CuEquivariance OpenEquivariance OpenEquivariance e3j (CUDA) e3j (Pallas GPU) e3nn (f32)† CuEquivariance OpenEquivariance OpenEquivariance e3j (CUDA) e3j (Pallas GPU)
✓ ✓
✓ ✓
63 ± 1 1245 ± 3 767 ± 7 806 ± 7 446 ± 1 –
65 ± 1 1287 ± 2 627 ± 4 598 ± 3 709 ± 2 1267 ±15
64 ± 0 1263 ±18 574 ± 3 504 ± 2 638 ± 5 1292 ± 9
62 ± 2 1264 ±22 573 ± 4 233 ± 1 560 ± 3 1319 ±10
60 ± 1 1248 ±15 564 ± 4 205 ± 1 623 ± 4 1328 ±30
56 ± 0 1059 ± 8 451 ± 2 470 ± 2 310 ± 0 –
53 ± 1 1104 ± 3 458 ± 1 190 ± 0 423 ± 2 1407 ± 3
52 ± 1 1119 ±12 411 ± 9 185 ± 1 402 ± 2 1432 ±18
51 ± 1 1118 ±22 324 ± 6 151 ± 0 317 ± 1 1483 ±22
27 ± 1 1128 ±29 253 ± 2 154 ± 1 183 ± 0 1599 ±12
7
Preprint
We also compare the performance of our CUDA and Pallas GPU kernels against the SO(2) convolution [23]. Our convolution kernels are 4 to 6 times faster than a JAX-based SO(2) convolution up to L = 3 (see [14] for implementation details), while curves cross at L = 4 for the forward, and L = 5 for the backward (see figure 7) when force-collapsing the multiplicity to maintain an equivalent level of expressivity. Further details and plots can be found in Appendix B.3. It is worth noting that the authors of UMA [5], also produced a Triton based kernel optimization for the convolution which reduces runtime by a third though fixed for L = 2. Assuming this improvement was ported into the JAX version it would still remain ∼ 6 times slower than our Pallas GPU convolution kernel. 3.2
E ND - TO - END BENCHMARKS
In order to assess the library in a practical setting, we performed end-to-end MLIP profiling and benchmarks of MACE and NequIP architectures. We compare our implementation with the e3nn_jax baseline of the mlip library [14], chosen for its unified and fast integration in downstream workflows (batched inference and simulations or relaxations), and to make sure that the backends are directly comparable. For each architecture, we connected e3j, cuEquivariance and OpenEquivariance as numerically equivalent message-passing backends. Although MLIP models have linear O(N ) complexity in the total number of atoms N , the messagepassing step scales with the number of edges and the average connectivity of the graph is typically larger than 40. For our MACE model (a) variant, message materialization may effectively bound achievable system sizes to around 13 thousand atoms before reaching memory overflows, while fused message-passing makes it possible to process up to 64 thousand atoms on a single NVIDIA® H100 GPU with 80 GB of global memory. Our results for MACE and NequIP are detailed in figures 3 and 4. It is worth noting that while not tested in this paper, the library can also be deployed across a number of alternative MLIP architectures such as GRACE [29] and Equiformer / E2Former [25, 27, 30, 31].
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). Backend e3j (CUDA) OpenEquivariance
Deterministic
Open-source
✓ ✓
cuEquivariance e3j (Pallas GPU) OpenEquivariance e3nn
ms/step
ns/day
✓ ✓
8.821 ± 0.045 16.137 ± 1.063
9.80 ± 0.05 5.38 ± 0.35
✓ ✓ ✓
6.586 ± 0.013 4.906 ± 0.014 10.041 ± 0.040 68.298 ± 0.063
13.12 ± 0.03 17.61 ± 0.05 8.60 ± 0.03 1.27 ± 0.00
One point to note is that in all MACE benchmarks, the symmetric contraction relies on the CUDA tensor product kernel of e3j with channel-mixing mode MAP. This helps enforce numerical consistency and allowed us to run stable simulations from a single trained checkpoint. It also has a relatively small impact with the correlation = 2 results presented here. See Appendix C.1 for experiments at correlation = 3 involving the proprietary cuEquivariance kernel for the symmetric contraction step and more discussion. Paradoxically, table 4 shows OpenEquivariance leads to slower simulations with the deterministic convolution kernel (see figure 2). This gap increases dramatically with the number of padding edges (above 30 ms/step with 25% padding), indicating their kernel hangs waiting for the slowest block accumulating messages on the single padding node. Our kernels flag padding edges so that work on these edges can be skipped, avoiding this overhead. 8
Preprint
MACE
NequIP
400 750
Runtime (ms)
TPU v6e
320
600
240
Backend
450
e3j (Pallas TPU) e3nn
160
300
80
0
150
0
2500
5000 7500 Number of atoms
10000
0
12500
0
1000
2000 3000 Number of atoms
4000
5000
Figure 3: Runtime (ms) for end-to-end force inference of MACE and NequIP on TPU. Inference is performed through integration of additional convolution backends for the mlip library [14] and run on real protein systems with 5Å cutoff. Hyperparameters can be found in table 3: MACE (a) (correlation 2) has 2 layers and NequIP has 5 layers, both have 128 channels. MACE
NequIP 500
200
Backend
Runtime (ms)
GPU H100
400
e3j (CUDA)
150
e3j (Pallas GPU) † 300
e3nn † cuEquivariance †
100
OpenEquivariance
200
OpenEquivariance † 50
100
0
†
non-deterministic
0 0
2000
4000 6000 8000 Number of atoms
10000
12000
0
2000
4000 6000 8000 Number of atoms
10000
12000
Figure 4: Runtime (ms) for end-to-end force inference of MACE and NequIP on GPU. Inference is performed through integration of additional convolution backends for the mlip library [14] and run on real protein systems with 5Å cutoff. Hyperparameters can be found in table 3: MACE (a) (correlation 2) has 2 layers and NequIP has 5 layers, both have 128 channels.
4
D ISCUSSION
Equivariant architectures have often been criticized for being computationally heavy, with learned equivariance often being put forward as an efficient inference time alternative. Multiple architectures have recently moved away from the strict inductive bias in favor of transformer architectures [32, 33], using the advantageous engineering of the transformer architecture as motivation for the change. In this paper we show that with proper engineering, full Clebsch-Gordan tensor product and associated message passing convolutions can be efficiently implemented on both GPUs and TPUs, reaching over 80% of HBM throughput. This should provide a more even playing field when comparing highly engineered transformer architectures with equivariant networks. It is worth noting that while dedicated kernels do improve the efficiency of the Clebsch-Gordan coefficients, and appear to be the best performing operation for L < 5, they remain at a disadvantage in terms of scaling in L compared to other methods, in particular the SO(2) convolution. Full treatment of the tensor product however has the benefit of being applicable to arbitrary feature vector pairs, and not solely in combination with harmonic features. A fair comparison would also require for SO(2) convolutions to also receive dedicated engineering, an effort that has begun with the second version of UMA [5] but that remains poorly explored by the community. One key point to note however, is that the achieved throughput remains below maximum bandwidth of the H100 GPU and v6e TPU, suggesting further improvements may be possible. We believe the release of e3j as an open-source library will prove a valuable new starting point for the community to continue optimizing these operations on current and future hardware. 9
Preprint
AI USE STATEMENT In this work, we used generative AI tools to help implement methods. We have not used generative AI tools to help develop theoretical models or conceptual frameworks, formulate mathematical claims, propose or refine hypotheses, design or provide feedback on research methodology or experiments, assist with translation, support qualitative and thematic data analysis, interpret results, and the rest of the required disclosure tasks (generate synthetic data sets, help develop theoretical models or conceptual frameworks, provide critical ingredients for proving mathematical claims, assist in the writing of proofs, clean and reformat dataset) are not applicable to this work. Additionally, we used generative AI tools to modify scientific figures and tables, create or edit software code. We have reviewed all AI-assisted work. LLM-generated code was reviewed by more than 2 authors and tested for correctness. We take responsibility for the final content of this work, including text, claims or artifacts produced with the aid of generative AI.
ACKNOWLEDGMENTS This work was supported by Cloud TPUs from Google’s TPU Research Cloud (TRC). We would like to express our gratitude to Sébastien B. and Oliver B. for encouragement and support in the early development of the library, and to warmly thank Marco C. for his continued assistance in the MLIP integration effort.
R EFERENCES [1] Nathaniel Thomas, Tess Smidt, Steven Kearnes, Lusann Yang, Li Li, Kai Kohlhoff, and Patrick Riley. Tensor field networks: Rotation- and translation-equivariant neural networks for 3d point clouds, 2018. URL https://arxiv.org/abs/1802.08219. [2] Simon Batzner, Albert Musaelian, Lixin Sun, Mario Geiger, Jonathan P. Mailoa, Mordechai Kornbluth, Nicola Molinari, Tess E. Smidt, and Boris Kozinsky. E(3)-equivariant graph neural networks for data-efficient and accurate interatomic potentials. Nature Communications, 13(1), May 2022. ISSN 2041-1723. doi: 10.1038/s41467-022-29939-5. URL http://dx.doi. org/10.1038/s41467-022-29939-5. [3] Ilyes Batatia, Dávid Péter Kovács, Gregor N. C. Simm, Christoph Ortner, and Gábor Csányi. MACE: Higher Order Equivariant Message Passing Neural Networks for Fast and Accurate Force Fields, 2023. URL https://arxiv.org/abs/2206.07697. [4] Dávid Péter Kovács, J. Harry Moore, Nicholas J. Browning, Ilyes Batatia, Joshua T. Horton, Yixuan Pu, Venkat Kapil, William C. Witt, Ioan-Bogdan Magdău, Daniel J. Cole, and Gábor Csányi. MACE-OFF: Transferable Short Range Machine Learning Force Fields for Organic Molecules, 2025. URL https://arxiv.org/abs/2312.15211. [5] Brandon M. Wood, Misko Dzamba, Xiang Fu, Meng Gao, Muhammed Shuaibi, Luis BarrosoLuque, Kareem Abdelmaqsoud, Vahe Gharakhanyan, John R. Kitchin, Daniel S. Levine, Kyle Michel, Anuroop Sriram, Taco Cohen, Abhishek Das, Ammar Rizvi, Sushree Jagriti Sahoo, Zachary W. Ulissi, and C. Lawrence Zitnick. Uma: A family of universal models for atoms, 2026. URL https://arxiv.org/abs/2506.23971. [6] Grzegorz Kaszuba, Tomasz Krakowski, Bartosz Ziegler, Andrzej Jaszkiewicz, and Piotr Sankowski. Implicit modeling of equivariant tensor basis with Euclidean turbulence closure neural network. Physics of Fluids, 37, 02 2025. doi: 10.1063/5.0249490. [7] Varun Shankar, Shivam Barwey, Zico Kolter, Romit Maulik, and Venkatasubramanian Viswanathan. Importance of equivariant and invariant symmetries for fluid flow modeling, 2023. URL https://arxiv.org/abs/2307.05486. [8] Oliver T. Unke, Mihail Bogojeski, Michael Gastegger, Mario Geiger, Tess Smidt, and KlausRobert Müller. Se(3)-equivariant prediction of molecular wavefunctions and electronic densities, 2021. URL https://arxiv.org/abs/2106.02347. 10
Preprint
[9] James Bradbury, Roy Frostig, Peter Hawkins, Matthew James Johnson, Chris Leary, Dougal Maclaurin, George Necula, Adam Paszke, Jake VanderPlas, Skye Wanderman-Milne, and Qiao Zhang. JAX: composable transformations of Python+NumPy programs, 2018. URL http://github.com/jax-ml/jax. [10] Adam Paszke, Sam Gross, Soumith Chintala, Gregory Chanan, Edward Yang, Zachary DeVito, Zeming Lin, Alban Desmaison, Luca Antiga, and Adam Lerer. Automatic differentiation in pytorch. In NIPS-W, 2017. [11] NVIDIA Corporation. CUDA C++ Programming Guide, 2025. URL https://docs. nvidia.com/cuda/archive/12.8.1/cuda-c-programming-guide/index. html. [12] Mario Geiger, Tess Smidt, Alby M., Benjamin Kurt Miller, Wouter Boomsma, Bradley Dice, Kostiantyn Lapchevskyi, Maurice Weiler, Michał Tyszkiewicz, Simon Batzner, Dylan Madisetti, Martin Uhrin, Jes Frellsen, Nuri Jung, Sophia Sanborn, Mingjian Wen, Josh Rackers, Marcel Rød, and Michael Bailey. Euclidean neural networks: e3nn, April 2022. URL https: //doi.org/10.5281/zenodo.6459381. [13] Mario Geiger and Tess Smidt. e3nn: Euclidean neural networks, 2022. URL https:// arxiv.org/abs/2207.09453. [14] Christoph Brunken, Olivier Peltre, Heloise Chomet, Lucien Walewski, Manus McAuliffe, Valentin Heyraud, Solal Attias, Martin Maarand, Yessine Khanfir, Edan Toledo, Fabio Falcioni, Marie Bluntzer, Silvia Acosta-Gutiérrez, and Jules Tilly. Machine learning interatomic potentials: library for efficient training, model development and simulation of molecular systems, 2025. URL https://arxiv.org/abs/2505.22397. [15] Oliver T. Unke and Hartmut Maennel. E3x: E(3)-equivariant deep learning made easy. arXiv preprint arXiv:2401.07595, 2024. [16] Hartmut Maennel, Oliver T. Unke, and Klaus-Robert Müller. Complete and efficient covariants for 3d point configurations with application to learning molecular quantum properties, 2024. URL https://arxiv.org/abs/2409.02730. [17] Vivek Bharadwaj, Austin Glover, Aydin Buluc, and James Demmel. An efficient sparse kernel generator for o(3)-equivariant deep networks, 2025. URL https://arxiv.org/abs/ 2501.13986. [18] YuQing Xie, Ameya Daigavane, Mit Kotak, and Tess Smidt. The price of freedom: Exploring expressivity and runtime tradeoffs in equivariant tensor products. In Forty-second International Conference on Machine Learning, 2025. URL https://openreview.net/forum?id= EvIwwGYTLc. [19] Shengjie Luo, Tianlang Chen, and Aditi S. Krishnapriyan. Enabling efficient equivariant operations in the fourier basis via gaunt tensor products, 2024. URL https://arxiv.org/ abs/2401.10216. [20] YuQing Xie, Ameya Daigavane, Mit Kotak, and Tess Smidt. Asymptotically fast clebsch-gordan tensor products with vector spherical harmonics, 2026. URL https://arxiv.org/abs/ 2602.21466. [21] Valentin Heyraud, Zachary Weller-Davies, and Jules Tilly. Integral formulas for vector spherical tensor products, 2026. URL https://arxiv.org/abs/2603.08630. [22] Anton Bochkarev, Yury Lysogorskiy, and Ralf Drautz. Fast contracted clebsch–gordan tensor products for equivariant graph neural networks, 2026. URL https://arxiv.org/abs/ 2605.15073. [23] Saro Passaro and C. Lawrence Zitnick. Reducing SO(3) Convolutions to SO(2) for Efficient Equivariant GNNs. In Proceedings of the 40th International Conference on Machine Learning, ICML’23. JMLR.org, 2023. 11
Preprint
[24] Xiang Fu, Brandon M. Wood, Luis Barroso-Luque, Daniel S. Levine, Meng Gao, Misko Dzamba, and C. Lawrence Zitnick. Learning smooth and expressive interatomic potentials for physical property prediction, 2025. URL https://arxiv.org/abs/2502.12147. [25] Yunyang Li, Lin Huang, Zhihao Ding, Xinran Wei, Chu Wang, Han Yang, Zun Wang, Chang Liu, Yu Shi, Peiran Jin, Tao Qin, Mark Gerstein, and Jia Zhang. E2former: An efficient and equivariant transformer with linear-scaling tensor products. In The Thirty-ninth Annual Conference on Neural Information Processing Systems, 2025. URL https://openreview. net/forum?id=ls5L4IMEwt. [26] Lin Huang, Chengxiang Huang, Ziang Wang, Yiyue Du, Chu Wang, Haocheng Lu, Yunyang Li, Xiaoli Liu, Arthur Jiang, and Jia Zhang. E2former-v2: On-the-fly equivariant attention with linear activation memory, 2026. URL https://arxiv.org/abs/2601.16622. [27] Yi-Lun Liao, Brandon M Wood, Abhishek Das, and Tess Smidt. Equiformerv2: Improved equivariant transformer for scaling to higher-degree representations. In The Twelfth International Conference on Learning Representations, 2024. URL https://openreview.net/ forum?id=mCOBKZmrzD. [28] Jonathan Godwin, Thomas Keck, Peter Battaglia, Victor Bapst, Thomas Kipf, Yujia Li, Kimberly Stachenfeld, Petar Veličković, and Alvaro Sanchez-Gonzalez. Jraph: A library for graph neural networks in jax., 2020. URL http://github.com/deepmind/jraph. [29] Yury Lysogorskiy, Anton Bochkarev, and Ralf Drautz. Graph atomic cluster expansion for foundational machine learning interatomic potentials. npj Computational Materials, 12(1), February 2026. ISSN 2057-3960. doi: 10.1038/s41524-026-01979-1. URL http://dx. doi.org/10.1038/s41524-026-01979-1. [30] Yi-Lun Liao and Tess Smidt. Equiformer: Equivariant graph attention transformer for 3d atomistic graphs. In The Eleventh International Conference on Learning Representations, 2023. URL https://openreview.net/forum?id=KwmPfARgOTD. [31] Yi-Lun Liao, Alexander J. Hoffman, Sabrina C. Shen, Alexandre Duval, Sam Walton Norwood, and Tess Smidt. Equiformerv3: Scaling efficient, expressive, and general se(3)-equivariant graph attention transformers, 2026. URL https://arxiv.org/abs/2604.09130. [32] Eric Qu, Brandon M. Wood, Aditi S. Krishnapriyan, and Zachary W. Ulissi. A recipe for scalable attention-based mlips: unlocking long-range accuracy with all-to-all node attention, 2026. URL https://arxiv.org/abs/2603.06567. [33] Ahmed A. Elhag, Arun Raja, Alex Morehead, Samuel M. Blau, Hongtao Zhao, Christian Tyrchan, Eva Nittinger, Garrett M. Morris, and Michael M. Bronstein. Learning inter-atomic potentials without explicit equivariance, 2026. URL https://arxiv.org/abs/2510. 00027. [34] Eugene Wigner. Group Theory And Its Application to the Quantum Mechanics of Atomic Spectra. Elsevier Science, 1959. [35] Brian C. Hall. An Elementary Introduction to Groups and Representations, 2000. URL https://arxiv.org/abs/math-ph/0005032. [36] Robert S. Womersley. Efficient Spherical Designs with Good Geometric Properties, page 1243–1285. Springer International Publishing, 2018. ISBN 9783319724560. doi: 10.1007/978-3-319-72456-0_57. URL http://dx.doi.org/10.1007/ 978-3-319-72456-0_57. [37] Norman P. Jouppi, Cliff Young, Nishant Patil, David Patterson, Gaurav Agrawal, Raminder Bajwa, Sarah Bates, Suresh Bhatia, Nan Boden, Al Borchers, Rick Boyle, Pierre-luc Cantin, Clifford Chao, Chris Clark, Jeremy Coriell, Mike Daley, Matt Dau, Jeffrey Dean, Ben Gelb, Tara Vazir Ghaemmaghami, Rajendra Gottipati, William Gulland, Robert Hagmann, C. Richard Ho, Doug Hogberg, John Hu, Robert Hundt, Dan Hurt, Julian Ibarz, Aaron Jaffey, Alek Jaworski, Alexander Kaplan, Harshit Khaitan, Daniel Killebrew, Andy Koch, Naveen Kumar, Steve Lacy, James Laudon, James Law, Diemthu Le, Chris Leary, Zhuyuan Liu, Kyle Lucke, Alan Lundin, 12
Preprint
Gordon MacKean, Adriana Maggiore, Maire Mahony, Kieran Miller, Rahul Nagarajan, Ravi Narayanaswami, Ray Ni, Kathy Nix, Thomas Norrie, Mark Omernick, Narayana Penukonda, Andy Phelps, Jonathan Ross, Matt Ross, Amir Salek, Emad Samadiani, Chris Severn, Gregory Sizikov, Matthew Snelham, Jed Souter, Dan Steinberg, Andy Swing, Mercedes Tan, Gregory Thorson, Bo Tian, Horia Toma, Erick Tuttle, Vijay Vasudevan, Richard Walter, Walter Wang, Eric Wilcox, and Doe Hyun Yoon. In-datacenter performance analysis of a tensor processing unit. ACM SIGARCH Computer Architecture News, 45(2):1–12, 2017. ISSN 0163-5964. doi: 10. 1145/3140659.3080246. URL http://dx.doi.org/10.1145/3140659.3080246.
A
M ATHEMATICAL BACKGROUND
In this appendix, we provide further theoretical details on the core components of e3j. Representation theory of the rotation and Euclid groups had ground-breaking applications in Quantum Mechanics, where they notably provided a first derivation of the Hydrogen energy levels and their degeneracy, see e.g. Wigner [34]. For more contemporary introductions to the subject, readers may refer to [15, 35]. A.1
E UCLIDEAN EQUIVARIANCE
The Euclid group E3 = O3 ⋉ R3 describes the possible changes of frames of reference over the Euclidean 3-space, i.e. the compositions of translations, rotations and reflections. Euclidean equivariance is the property by which a function on the Euclidean space transforms consistently with its inputs upon any Euclidean transform g ∈ E(3), for instance, with a force function F(r, z): F(g · r, z) = g · F(r, z). n×3
(1)
n
where r ∈ R denotes a matrix of atomic positions, z ∈ N denotes a vector of atomic numbers, and n is the number of atoms of the system or region of interest. In general, the Euclidean group may not only act on (n copies of) R3 , but also on general (real or complex) vector spaces called representations of E3 (also called E3 -modules): they consist of pairs (V, ρ) where the vector space V is equipped with a smooth group morphism ρ mapping any Euclidean transform g ∈ E3 to an invertible matrix ρg ∈ GL(V ). A function F : V → V ′ is called equivariant if the following diagram is commutative: V
F
ρ′g
ρg
V
V′
F
(2)
V′
While morphisms of E3 -representations are usually assumed linear, note that the above definition of equivariance applies to non-linear functionals just as well, in particular polynomial functionals which are of particular importance in the classification of E3 -representations. A fundamental result is that any orthogonal representation V of O3 can be decomposed into a direct sum of irreducible representations or irreps, each irrep being chosen among a well known classification of possible fundamental types (related to harmonic polynomials over R3 ). Note that the dimensions of irreducible representations translate into tangible observables in quantum chemistry (QC), where l and 2l + 1 respectively determine the symmetry (S, P, D...) and degeneracy (half the maximal number of occupying electrons) of an energy level in a hydrogen-like atom. A.2
H ARMONIC POLYNOMIALS
Definition. The representation theory of SO3 is closely related to the harmonic polynomials of C[x, y, z]. Homogeneous, degree-l polynomials are finite-dimensional vector spaces, naturally equipped with an SO3 action. Because the Laplacian operator ∆ (trace of the hessian) is also invariant under SO3 , the space Yl ⊂ Cl [x, y, z] of degree-l harmonic polynomials (satisfying ∆P = 0) is a sub-representation, i.e. a subspace stable under SO3 . ∆Ylm =
∂ 2 Ylm ∂ 2 Ylm ∂ 2 Ylm + + =0 ∂x2 ∂y 2 ∂z 2 13
Preprint
Any irreducible representation of SO3 (i.e. a "smallest" vector space with an SO3 -action, having no other stable strict subspace than 0) is isomorphic to some Yl , of odd dimension 2l + 1. L Any larger representation of SO3 can be decomposed as a direct sum of irreducible representations l (Yl )kl . Equivariance. For every degree l, the maps Yl : R3 → Yl ≃ C2l+1 are equivariant non-linear embeddings, typically used as a basis for learning more complex non-linear equivariant representations fθ : V → V ′ within e.g. deep MLIP networks. This means that for every rotation g ∈ SO3 , one may construct the so-called Wigner D-matrix Dg , acting on the 2l + 1 space of polynomial activations Yl in the following commutative diagram: R3
Yl
g
R3
Yl Dg
Yl
(3)
Yl
The spaces Yl consist of the basic pieces of any Euclidean representation V , as any irreducible representation of the group of rotations O3 is isomorphic to some Yl for some l. Additionally, irreducible E3 representations carry a parity label ± (even/odd) dictating whether reflections act with a sign or as the identity. See Appendix G for further mathematical details. A.3
K EY O PERATIONS IN S COPE
Equivariant architectures typically enforce the equivariance constraint (1) by a few common design patterns and building blocks, the e3j package provides a harmonized and functional API around these: • Harmonics: Restrict available geometric information to the edge vectors (rab ) ∈ RnE ×3 . This already enforces translation invariance. Then expand edge vectors rab ∈ R3 with harmonic embeddings Yl (rab ) where Yl denotes one bank of 2l + 1 rotation-equivariant activation filters built from the 2l + 1 degree-l harmonic polynomials, typically concatenated over degrees l = 0, . . . , lmax with lmax rarely exceeding 3: Yl (r) = Ylm (r) | m = −l . . . l (4) • Linear mixing: Rescale or mix channels and multiplicities linearly between irreducible features of a same degree l. While rarely a bottleneck by itself, there is opportunity for scalar mixings to be fused e.g. within message-passing operations, where edge scalars representing chemical species and radial embeddings of interatomic distances are coupled with rotation-equivariant features. X lm′ ,k′ LW (x)lm,k = Wkkm (5) ′ m′ x k′ m′
• Tensor product: Couple latent equivariant features with Clebsch-Gordan tensor products z = x ⊗C y, where x may for instance denote latent node features from the previous layer and y the harmonic embedding of an edge vector, or a more general latent feature vectors organized in irreps. When x and y are irreducibles of degree l and l′ , their tensor product z is obtained from the 2 min(l, l′ )+1 bilinear pairings (or "paths") of output degrees L = |l − l′ |, . . . , l + l′ , given by: X LM lm l′ m′ (x ⊗C y)LM = Clm,l y (6) ′ m′ x m+m′ =M LM where Clm,l ′ m′ are the Clebsch-Gordan coefficients.
• Message passing: Aggregate the edge-wise features, usually computed through combination of harmonics projection, tensor product and linear or scalar mixing. The message passing aggregation tends to be the operational bottleneck once efficient tensor product operations 14
Preprint
are implemented. A typical message passing layer in the context of equivariant GNNs can be written as: X m′b = Lsab (xa ⊗C Y(rab )) (7) a∈N (b)
L
where Y(rab ) = l Yl (rab ) and Lsab denotes a linear mixing with radial edge scalars sab , which typically depend on interatomic distances ||rab || via a radial basis function (RBF) embedding followed by a multi-layer perceptron (MLP). While we have used MLIP as a testing ground for realistic workloads, we have not yet included operations specifically targeted at particular MLIP models. The main example is the Symmetric Contraction module in MACE. We did however include benchmarks with the existing cuEquivariance Symmetric Contraction kernel in appendix C.1.
B
C OMPARISON WITH C UBICALLY S CALING T ENSOR P RODUCTS
B.1
G AUNT AND V ECTOR S IGNAL T ENSOR P RODUCTS
In this section, we benchmark the e3j tensor product against the Gaunt tensor product (GTP) [19] and the vector signal tensor product (VSTP) [20–22], which evaluate the tensor product as an integral over the sphere rather than a contraction over Clebsch-Gordan coefficients. The GTP and VSTP are complementary: the GTP covers only the symmetric paths, where the sum l1 + l2 + l3 of all of the irreps entering into the tensor product is even, while the VSTP covers only the skew-symmetric paths where the sum is odd. Forward
Backward
Runtime (μs)
Symmetric
105
104
103
GTP VSTP e3j (CUDA)
Runtime (μs)
Skew-symmetric
10
5
e3j (Pallas GPU)
104
103
2
3
4
5 ℓmax
6
7
8
2
3
4
5 ℓmax
6
7
8
Figure 5: Comparison with Gaunt and Vector Signal Tensor Products. Benchmarks show the runtime scaling with L = lmax with a fixed batch size B = 32,768 and channels C = 128. Every Clebsch-Gordan path (l1 , l2 , l3 ) is symmetric or skew-symmetric, according to the parity of l1 + l2 + l3 . The symmetric panel, where the sum is even, benchmarks the e3j tensor product on symmetric paths against the Gaunt tensor product (GTP) [19], whilst the skew-symmetric panel benchmarks against the vector signal tensor product (VSTP) [20, 21]. Both the GTP and the VSTP are evaluated through integral formulas on a spherical t-design quadrature [36]. Runtimes are obtained on a single NVIDIA H100. Efficient implementations of both the VSTP and GTP rely on the tensor product emitting a single copy of each output degree rather than one per path, and, relatedly, that a weighted tensor product has weights taking a factorised form wll13l2 = al1 bl2 cl3 . In this benchmark, we therefore collapse the e3j 15
Preprint
output by degree to ensure the output spaces match in all cases. We also choose to time the tensor product, since factorised weights act as linear maps on the input and output spaces that are identical in all three cases. In practice, we compute the GTP and VSTP from their integral formulas, which are evaluated on a spherical t-design quadrature [36]. To match the paths of the GTP and VSTP, we utilize the parity selection rule p1 p2 = p3 that the Clebsch-Gordan product already enforces. With natural-parity irreps pl = (−1)l for both the inputs and the targets, the parity rule enforces (−1)l1 +l2 +l3 = 1 and keeps exactly the symmetric paths. Conversely, with the anti-natural pl = (−1)l+1 irreps, only the skew-symmetric paths are allowed. The results are shown in Figure 5. For L ≤ 3, where the cost of the quadrature operations dominates, the e3j kernel is 2 to 4 times faster than the GTP and 4 to 6 times faster than the VSTP. As L grows, the theoretical scaling in L becomes relevant (O(L4 ) for quadrature methods vs O(L5 ) for sparse Clebsch-Gordan contraction). The forward pass curves cross at L = 6 and L = 7 for the GTP/VSTP respectively, and at L = 5 and L = 6 for the backward, showing that e3j remains faster in the regime practical for MLIP applications. It is worth noting that there are also other implementations of the GTP and VSTP that have better theoretical scaling in L [20], but that we find to be slower in practice over this range of degrees. B.2
M ATRIX T ENSOR P RODUCTS FROM e3x
Another tensor product operation with an efficient O(L4 ) scaling has been proposed by the authors of the e3x library [15, 16]. We follow Xie et al. [18] and refer to this operation as Matrix Tensor Product (MTP). The MTP is implemented in the FusedTensor module of e3x. This operation computes the tensor product between feature vectors in three steps. First, considering the operation couples two input vectors containing irreps features of order 0 ≤ l ≤ L, each input vector is mapped to a (2˜l + 1) × (2˜l + 1) square matrix, with ˜l = ⌈L/2⌉. This map corresponds to the isomorphism L M
Yl ∼ = Yl̃ ⊗ Yl̃ ,
(8)
l=0
where elements of the right hand-side space can be seen as square matrices produced by an outer product of two feature vectors in Yl̃ . Then, the square matrices encoding the inputs are multiplied. Finally, the resulting square matrix is mapped back to the output feature vector. The matrix product exhibits an efficient O(L3 ) scaling, but the conversions between square matrices and feature vectors yield an overall scaling of O(L4 ). In this section, we benchmark the e3j tensor product against the MTP. Note that for a given irrep path, the MTP is proportional to the GTP-VSTP operations, although the proportionality coefficient sometimes vanishes, so that the MTP is in general strictly less expressive than the GTP-VSTP. We compare the FusedTensor module of e3j against the same e3j tensor product used in the benchmark of Appendix B.1. As the FusedTensor implementation includes the path weights wll13l2 , we include linear mixings on both the inputs and the output feature vectors of the e3j tensor product, so that both operations are strictly equivalent. Note that these additional linear mixings account for its higher runtime compared to the version of Appendix B.1. The results are shown in Figure 6. For L ≤ 8 the e3j tensor product is faster than the e3x implementation. As L grows, the efficient O(L4 ) scaling of the MTP reduces the gap with the e3j implementation, though e3j remains faster on this domain. B.3
SO2 C ONVOLUTION
In order to compare with SO2 convolution fairly we evaluated the e3j convolution implementations with a set of coefficients that collapses output multiplicities, i.e. sums all isomorphic copies of a given irreducible output space together. The reduced output feature dimension means a cost on expressivity, at the benefit of a smaller memory traffic. Note that while e3j convolution kernels natively support any set of coefficients, they have not been optimized for these smaller problem shapes. We however notice that e3j is between 4x and 6x faster at ℓmax ≤ 3, while matching SO2 convolution at ℓmax = 4, see figure 7. Although the XLA compilation of a plain JAX implementation of SO2 convolution already proves very efficient, a 16
Preprint
Runtime (μs)
Forward
Backward
Library
104
e3j (CUDA) e3j (Pallas GPU) e3x
103 2
3
4
5 ℓmax
6
7
8
2
3
4
5 ℓmax
6
7
8
Figure 6: Comparison of e3j and e3x’s FusedTensor tensor products. Benchmarks show the runtime scaling with L = lmax with a fixed batch size B = 32,768 and channels C = 128. Every Clebsch-Gordan path (l1 , l2 , l3 ) is symmetric. Runtimes are obtained on a single NVIDIA H100. faithful comparison of the method would require similar low-level engineering on the SO2 convolution implementation. Forward
Backward
Runtime (μs)
105
104
SO2 e3j (CUDA) e3j (unfused) e3j (Pallas GPU)
103
1
2
3
4
5
6
1
2
ℓmax
3 ℓmax
4
5
Figure 7: Comparison of SO2 convolution with Clebsch-Gordan convolution. Benchmarks show the runtime scaling with L = lmax with fixed number of nodes N = 2048, number of edges Ne = 45 × N , and fixed number of channels C = 128. In contrast to other convolution benchmarks with third-party baselines, output multiplicities are collapsed to one by providing an ad-hoc set of sparse COO coefficients to match a JAX implementation of the SO2 convolution algorithm described in [23]. Runtimes are obtained on a single NVIDIA H100.
C
A DDITIONAL E ND - TO - END B ENCHMARKS
C.1
S YMMETRIC CONTRACTION BENCHMARKS
Most of the MACE model benchmarks reported in the main text use the same implementation for the SymmetricContraction operation, which consists of: • a PowerExpansion of node features (quadratic or cubic) using the tensor product kernels of e3j with channel mixing-mode MAP5 , • a LinearIndexwise projection of the concatenated higher-order features using weights that depend on the atomic species. While conceptually simple, this implementation is suboptimal as the power expansion generates large multiplicities that all collapse eventually after the linear projection steps. The GMEM materialization 5 Note that not all CUDA backends expose this functionality, and that the batch and channel axes are not contiguous in the optimal trailing channels layout.
17
Preprint
of the large intermediate array can be skipped by a dedicated SymmetricContraction fusing both operations, as reported in figure 8. MACE (a)
MACE (b) 50
50 40
Runtime (ms)
40
Convolution backend e3j (CUDA) e3j (Pallas GPU)
30
30
cuEquivariance Symmetric contraction
20
20
e3j cuEquivariance 10
10
0
0 0
2000
4000 6000 8000 Number of atoms
10000
12000
0
2000
4000 6000 8000 Number of atoms
10000
12000
Figure 8: Effect of the SymmetricContraction backend for the MACE model on GPU. We compare the effect of replacing the naive e3j-based power expansion, followed by a species-wise linear projection of higher-order features, with the dedicated SymmetricContraction kernel of cuEquivariance. The two variants (a) and (b) of the MACE model are detailed in table 3. An optimized implementation may follow the algorithm proposed in the original MACE paper [3], in device code. The plain JAX implementation of this algorithm however performs worse than the naive power expansion and linear projection, once an efficient CGTP is available. The e3nn forceinference benchmarks, illustrating the best efficiency one may reach with a pure JAX implementation, rely on this implementation. However, we cannot currently change the algorithm significantly without breaking numerical consistency. Restoring consistency requires a complex 1-to-1 transformation of learnable parameters and CG coefficient normalization choices to be resolved. All NPT simulation benchmarks use the same e3j-based SymmetricContraction. While isolating the effect of the convolution backend, this allows us to load a single model checkpoint to run stable NPT simulations. At the time of writing, e3j does not provide a dedicated SymmetricContraction kernel. Our current investigations seemed to show that inlining coefficients with a JIT compiled kernel (using Pallas or NVRTC) seems necessary to reach the same performance as cuEquivariance, given that our AOT compiled CUDA kernels reach due to the large number of coefficients and feature sizes that occur during this operation.
C.2
D ETERMINISM AND DEVIATIONS OF END - TO - END PREDICTIONS
Non-deterministic message aggregation may prove very efficient, since it a allows a kernel to directly loop and distribute work over edges (instead of nesting a potentially imbalanced loop over neighbors inside a loop over receiver nodes) and the memory cache hierarchy may efficiently hide the latency of memory-locked atomic operations. Valency imbalance for instance explains why the deterministic OpenEquivariance kernel leads to slower simulations than the non-deterministic, when enforcing a static edge count with all padding edges joining a single padding node, despite being faster in non-padded unique benchmarks. In some downstream workflows (relaxation, geometry optimization, ...), deterministic predictions may however be a important requirement. When relaxing a periodic cell (as in NPT simulations, see table 2) with a Monte-Carlo barostat, deterministic energy predictions are for instance a hard constraint: small energy fluctuations enter an exponential Boltzmann factor used for a Metropolis-Hastings rejection criterion, and non-deterministic message aggregation would lead to exploding simulations. In addition to the message-passing step, two sources of non-determinism or numerical noise may compound in the energy prediction: 18
Preprint
• addition of atomic energies E0 , which are multiple orders of magnitude larger than the geometry-induced energy variations, and may truncate the constant relative precision of floating-points data types, • aggregation of graph energies from node energy summands, which may scatter a very large number of contributions to a few scalar numbers with a high degree of concurrency and collisions. Both of these effects are analyzed in figure 9. In contrast, the force prediction may eliminate the addition of constants and replace non-deterministic scatter operations by deterministic gather operations in the computational graph, as illustrated by figure 10. We note that this mitigation of stochastic variations is not an automatic consequence of the differentiation process, since differentiating a mean-squared-error loss would presumably compound sources of non-determinism instead of eliminating most of them (through a product of the energy error with energy gradients). A rigorous analysis of the effect of non-determinism on model training behavior is however out of scope of the present work. Convolution backend
100
OpenEquivariance OpenEquivariance †
10
−1
100 10−2 10−1 10−3
Energy deviation (kT)
Energy deviation (eV)
101
cuEquivariance † e3j (CUDA) e3j (Pallas GPU) † †
non-deterministic
10−2 10−4 dense
scatter
dense + E0
scatter + E0
Energy aggregation
Figure 9: Run-to-run energy deviation of a MACE model (a) on GPU. The trained model used for NPT simulations (table 2) is evaluated on a batch of 8 water boxes totaling 21,016 atoms and 889,840 edges, over a 100 times. The experiment was repeated using different energy aggregation schemes: including or dropping the constant atomic energies (E0 ); using mlip’s default scatter aggregation of node energies over graphs or a dense matrix-vector product for the energy head. The plot represents 5th/95th percentiles as whiskers and 25th/75th percentiles as boxes. When a box is missing, it means that all runs were rigorously deterministic. Hyperparameters are detailed in table 3.
D
B ENCHMARKS DETAILS
D.1
T HROUGHPUT AND HBM
Our module-specific benchmarks mostly focus on kernel throughput, commonly defined as: sizeof(inputs) + sizeof(outputs) (9) runtime In addition to being asymptotically independent of I/O size, throughput is also bounded by the so-called global memory (GMEM) bandwidth, corresponding to the ideal throughput of an optimal array copy. Global memory is the long-lived data bank used by the device processors to load/store I/O data. It is also called high bandwidth memory (HBM) in manufacturer specifications, given its practical importance in delivering the best I/O throughput possible. throughput =
The NVIDIA® H100 graphical processing units (GPUs), on which most of our experiments were performed, advertises about 3.35 TB/s HBM. The Google®tensor processing units (TPUs) we could experiment with advertizes 1.20 TB/s HBM for v4 and 1.64 TB/s HBM for v6e (Trillium). Note constant technical progress is made on those characteristics, and a fair comparison should compare devices from the same year, and weigh those metrics by affordability. 19
Preprint
Convolution backend OpenEquivariance OpenEquivariance †
Force deviation (eV/Å)
10−4
cuEquivariance † e3j (CUDA)
10
e3j (Pallas GPU) †
−5
†
non-deterministic
10−6
10−7
dense
scatter
dense + E0
scatter + E0
Energy aggregation
Figure 10: Run-to-run force deviation of a MACE model (a) on GPU. The trained model used for NPT simulations (table 2) is evaluated in the same conditions as figure 9. Because XLA can eliminate the addition of atomic energies, and because the VJP of the scatter operation is a deterministic gather operation, we see that the addition of E0 and the aggregation method have no effect on the deviation of forces. Even with deterministic convolution kernels, forces have non-zero deviations due to some forward gather operations being transposed as non-deterministic scatter operations. Interestingly the Pallas GPU kernel, only deterministic in the forward pass, lies in between.
D.2
P RECISION
Although e3j also provides float64 binaries, all benchmarks were carried with I/O arrays in single float32 precision, which proves enough to run stable molecular dynamics simulations. Note however that JAX may internally resort to half-precision arithmetic in tensor contractions (matmul, einsum, . . . ) for faster execution, and caps to single-precision by default to avoid downsides incurred by undesired upcasts. On normal random input, our experiments show that all equivariance backends, regardless of platform, lead to similar accuracies of order 4 × 10−8 and 1 × 10−7 for tensor products and message-passing operations respectively, with respect to a common, deterministic float64 reference. The only exception is the e3nn backend which only yields about 4 × 10−4 elementwise accuracy when the environment variable JAX_DEFAULT_MATMUL_PRECISION is not set to highest. This significant gap in precision should therefore be considered when comparing e3nn with other backends in end-to-end benchmarks. D.3
C ONSIDERATIONS REGARDING PALLAS AND CUDA KERNELS
To yield speedups on the Google®tensor processing units (TPUs) which JAX targets as well, E3J also defines kernels written in Pallas, a domain-specific language (DSL) which is part of the JAX package and targets GPU and TPU compilation. While TPU benchmarks can only compare Pallas kernels of E3J with E3NN [12], the GPU benchmarks may compare CUDA and Pallas implementations of E3J with other CUDA implementations such as NVIDIA CuEquivariance (TM) and OpenEquivariance [17]. One advantage CUDA nonetheless brings over Pallas is the relative stability of compiled binaries and nvcc toolchain over JAX version dependencies. However, the Pallas language and Mosaic GPU compiler streamline the just-in-time (JIT) compilation of device code, enabling the production of very specialized and efficient kernels. In addition to problem shapes and sizes, that cannot be trivially defined from static parameters with ahead-of-time (AOT) compilation of a traditional CUDA kernel, JIT compilation from Python source lets one seamlessly inline the Clebsch-Gordan coefficients inside the produced assembly code. This can significantly reduce memory traffic from the unified SMEM/L1 cache. 20
Preprint
Note that other solutions such as NVIDIA’s NVRTC compiler enable JIT compilation of CUDA/C++source, a solution leveraged by OpenEquivariance. Our CUDA kernels do not make use of JIT compilation at this time, and stream through coefficients as an actual array buffer loaded from global memory. D.4
MLIP INTEGRATION
Hyperparameters defining the models used in the end-to-end benchmarks are detailed in table 3, they corresponding to the configuration fields of the mlip library [14]. While end-to-end MLIP integration gives the most significant results for applications, it is also a complex task that needs to be carried carefully in order to preserve numerical predictions and faithfulness of comparisons. Furthermore, while MLIPs consist of a primary motivation for e3j, we view the library as an all-purpose low-level tool whose usage may not be limited to the two particular architectures benchmarked in the present work. All the MLIP models compared were checked to match numerically regardless of the convolution or tensor product backend. Comparison with third-party implementations of models, or comparisons substituting additional blocks of the MACE model (such as SymmetricContraction) has not been performed at this time, due to the significant effort required to gain solid confidence in the faithfulness of final comparisons, and its orthogonality with the low-level engineering effort behind e3j.
MACE (a)
Table 3: Hyperparameters used in end-to-end MLIP benchmarks. MACE (b)
num_layers num_channels correlation node_symmetry l_max cutoff_angstrom num_rbf node_gating include_pseudotensors
2 128* 2 2 3 5 8 true false
num_layers num_channels correlation node_symmetry l_max cutoff_angstrom num_rbf node_gating include_pseudotensors
2 128* 3 1 3 5 8 true false
NequIP num_layers num_channels node_irreps l_max cutoff_angstrom num_rbf
E
5 32* 2x0e + 2x0o + 1o + 1e + 2e + 2o
2 5 8
T ENSOR P RODUCT AND M ESSAGE PASSING C ONVOLUTION KERNELS
In this appendix, we present algorithmic details of the three sets of kernels developed for e3j: CUDA, Pallas GPU, and Pallas TPU. The CUDA kernel offers the broadest applicability across various GPU architectures and is fully deterministic, while Pallas GPU focuses on performance for the latest GPU architectures and JAX versions. We first present the relevant notation and recall the operation that needs to be performed as part of the Clebsch–Gordan Tensor Product and associated message passing, and then present in order the CUDA, Pallas GPU and Pallas TPU kernels. In order to facilitate understanding for readers with different backgrounds, we have added a short section outlining the key hardware concepts for GPU and TPU in appendix F. 21
Preprint
E.1
M ATHEMATICAL D ETAILS
In this section we detail the notation used and provide a mathematical description for the operation performed by the kernel. For further details on the motivation for the construction of these operations from a theoretical standpoint, please refer to the original literature. Notation. • x: Left-hand-side features. In the Tensor Product, it is an arbitrary feature tensor, the Message Passing operation, it represents the sender node features. It is indexed along three axes: batch elements (usually nodes), equivariant features (see below for indexing notations), and channels • y: Right-hand-side features. In the Tensor Product, it is an arbitrary feature tensor, in the Message Passing operation, it usually represents (broadcasted) spherical harmonics embeddings. It is indexed along three axes: batch elements (usually edges), equivariant features (see below for indexing notations), and channels. • s: Edge scalars. Edge specific weighting, usually computed using a MLP on radial embedding projecting to the number of channels, for each edge. It is indexed along edges, equivariant features, and channels. • mb : Aggregated messages on receiver node. The output of the message passing convolution for each node. It is worth noting that it does not int principle, have the same feature dimension as x as multiplicities arise during the tensor product (these are in general contracted back to the number of channels in channel mixing). It is likewise indexed by nodes, equivariant features, and channels. • C: Clebsch–Gordan coefficients. Its values are indexed by outputs (i0), l.h.s input (i1) and r.h.s. input (i2) feature indexing (each i0, i1, i2 represents a ℓ, m combination) . • mab : Edge level message between a sender node and a receiver node. • a: Sender indexing • b: Receiver indexing • q: Channel indexing • p: sparse Clebsch–Gordan record index (when unrolled) • Nq : The number of channels • Np : The number of non-zero Clebsch–Gordan coefficients i0
• cp = Ci1pp,i2p : Coefficient value of record p Mathematical operations. Following the notation above, we can outline the operations performed by the kernels: • Tensor product of geometric features: This formula generalises to bilinear operation on arbitrary feature vectors, it includes the Clebsch–Gordan tensor products when appropriate weights are selected. Here we treat the sender / receiver indices as implicit as this operation can be performed arbitrarily on any batch element (i.e. nodes or edge features). z is an arbitrary notation to represents the output feature tensor following a tensor product. X i0 zi0,q = Ci1,i2 ∗ xi1,q ∗ yi2,q (10) i1,i2
• Message weighting with edge scalars: From the computed Tensor Product features on a given edge, one can weight the edge specific message before aggregation. This is usually done through a MLP mapping radial embeddings to a set number of channels. In the notation below, the index i3 maps the feature coordinate i0 to its irreducible block (piecewise-constant on irreducible subspaces), we write this dependency as i3(i0) for simplicity. The mapping from feature coordinates to scalar indices is performed at coefficient construction time from I/O representations. This operation can be viewed as a Tensor Product with the r.h.s. input being scalars instead of arbitrary geometric features. i3(i0),q
i0,q mi0,q ab = zab ∗ sab
22
(11)
Preprint
• Message Passing Convolution: The complete operation for the message passing convolution, per node, can be written concisely as this common particular version of (7): X mb = m̃ab where m̃ab = sab · (xa ⊗ yab ) (12) a∼b
given node features xa , edge features yab , scalar embeddings sab and letting the dot denote the scalar mixing operation. Expanding equations (10) and (11) yields XX i2,q i3(i0),q i0 mi0,q = Ci1,i2 xi1,q yab sab . (13) a b a∼b i1,i2
The equivalence between (12) and (13), reflecting the associativity of bilinear couplings, leads to different computation graphs and implementations, as depicted in figure 12. Differentiation. All of our kernels and associated primitives are infinitely differentiable. We outline below the mathematical derivation of their so-called reverse-mode AD primitives, or vector-jacobian products (VJP) rules, obtained by first differentiating the smooth function above a set of inputs, called primals, before transposing the linearized map. The transposed differential acts on cotangents (linear forms on tangent vectors) in the reversed direction, i.e. it maps output cotangents to primal cotangents so as to enable back-propagation of gradients in a complete neural network architecture. • Tensor product backward: By bilinearity of the tensor product, the Leibniz rule gives the forward-mode AD rule for z = x ⊗ y ∈ X ⊗ Y as: δz = δx ⊗ y + x ⊗ δy (14) where δx, δy denote tangent vectors of Tx X and Ty Y respectively, mapped to an output tangent δz ∈ Tz (X ⊗ Y ) by the linearized tensor product operation. Equation (14) thus defines a linear map Lx,y : Tx X ⊕ Ty Y → Tz (X ⊗ Y ) such that δz = Lx,y (δx, δy). The backward tensor product primitive is the adjoint (a.k.a. transpose)6 of Lx,y , which maps an output cotangent dz ∈ Tz∗ (X ⊗ Y ) to primal cotangents dx ∈ Tx∗ X and dy ∈ Ty∗ Y . In practice, this means transposing the partially applied tensor product maps x ⊗ − and − ⊗ y, respectively consuming δy and δx in (14). The transposition amounts to applying a permutation of indices. For dx, each record is permuted as (i0, i1, i2) 7→ (i1, i2, i0) and sorted by the new output index i1; for dy we use (i0, i1, i2) 7→ (i2, i1, i0) and sort by the new output index i2. The coefficient values are reordered identically. Writing z = tensor_product(Ci0 i1i2 , x, y) in reference to algorithm 1, the source cotangents can be evaluated with the same sparse kernel after permuting the coefficient records: dx = tensor_product(Ci1 i2i0 , y, dz) (15) dy = tensor_product(Ci2 i1i0 , x, dz) The operation is also illustrated in Fig. 11. Although this simple backward algorithm provides infinite differentiability, it has the disadvantage of loading output cotangents δz twice in on-device memory, which can be significantly expensive as the output feature dimensions scales as the product of input feature dimensions. • Message-passing backward: By duality, gather and scatter-add operations are swapped in reverse-mode differentiation and the backward message-passing rule can be expressed as a message-passing operation on the transposed graph, with multiple bilinear couplings being performed. Letting zab = xa ⊗ yab in (12), input cotangents may be computed from receiver cotangents dmb as: X dxa = sab · (yab ⊗ dmb ) b∼a dyab = sab · (xa ⊗ dmb ) conv_bwd(x, y, s, dm) = (16) +l X (l) (l,m) (l,m) zab · dmb dsab = m=−l
6 Formally, the adjoint L∗x,y is the composition or pull-back dz ◦ Lx,y . Composing the linear form dz with the linear map Lx,y defines a linear form on Tx X ⊕ Ty Y , or two linear forms on Tx X and Ty Y , a.k.a. cotangents.
23
Preprint
Because each cotangent is computed by an ad-hoc operation, (16) may hardly reuse device code from the forward pass yet specialized backward device code can prove more efficient7 . Instead of (16), if one views the message-passing as a trilinear coupling – see figure 12 and equation (12) – the VJP can be expressed as three distinct trilinear coupling calls to a single bigotimes routine, via a straightforward generalization of (15) to three operands. • Message-passing double backward: Deriving the second order rule from (16) is best viewed through the diagram 12 back-propagated twice, each differentiation of bigotimes yielding three transposed bigotimes calls writing to each of the input leaves. Denoting by δdx, δdy, δds the second-order variations of the primal inputs (first-order variations of cotangents returned by the backward pass), the associated variations δx, δy, δs, δdm returned by the backward rule are defined in terms of lower-order primitives as: δdm = conv(δdx, y, s) + conv(x, δdy, s) + conv(x, y, δds)
(
δx = δy x + δs x δy = δx y + δs y δs = δx s + δy s
( where
(17)
−, δx y, δx s = conv_bwd(δdx, y, s, dm) δy x, −, δy s = conv_bwd(x, δdy, s, dm) . (18) δs x, δs y, − = conv_bwd(x, y, δds, dm)
Note that defining AD rules via cyclic references to lower-order differentials, as (15) and (17-18) do, automatically yields infinitely differentiable JAX primitives. At this time e3j doesn’t define doublebackward kernels, in contrast with OpenEquivariance [17]. Double-backward kernels could save unused cotangents from being computed in (18), or save memory traffic by streaming through the primals in a single loop. Due to the increased number of I/O arrays, a GPU implementation would likely rely on implicit L1/L2 caching over shared memory buffering, yet these constraints do not apply on TPU whose VMEM slice is much larger. Double-backward optimization is left for future work. E.2
CUDA KERNELS
A significant design choice of e3j is to rely on a sparse and agnostic representation of total ClebschGordan coefficients, in which equivariant coordinates (traditionally a set of triplets (κ, ℓ, m) for the output and each of the two inputs, hence a total of 9 integers) are mapped to unfolded feature coordinates i0, i1, i2, for output, l.h.s input and r.h.s input, running over the whole I/O dimensions. Melding the equivariant coordinates together into opaque indices means we can rely on a very generic sparse tensor product algorithm, suitable for different applications. 3 The number of irreducible blocks in the sum scales as O(lmax ) – number of input degree pairs multiplied by the number of possibilities for the output degree – while the number of non-zero 2 coefficients within each block scales as lmax , due to the spin constraints M = m + m′ . Since the 2 5 feature dimension D scales as lmax , and our algorithm 1 scales as O(lmax ) = O(D5/2 ) in terms of floating point operations (FLOPs). In practical situations we targeted, lmax is a fixed hyperparameter of the MLIP model and the number of non-zero coefficients is a few hundreds or less.
E.2.1
CUDA T ENSOR P RODUCT OPERATION
The main bottleneck of the sparse, JAX-based implementation of the Tensor Product is the final scatter-reduction step, which scaled poorly to large input sizes. Our CUDA algorithm bypasses this bottleneck and implements a number of focused optimization, the key differentiating implementation factors of algorithm 1 are explained below: • Trailing channels: we avoid any inter-thread communication by sequentially reducing the feature axis, and parallelizing across the trailing channel dimension instead. The 32 threads of a warp (simultaneously scheduled to execute the same instruction) can thus process distinct channels in parallel, while doing the same work: i.e. warp lanes process distinct channels while traversing the same ordered sequence of CG coefficient records. 7 Our CUDA kernels prioritize code reuse while our Pallas kernels use specialized code in backward operations, see below.
24
Preprint
• Vectorization: processing 2 or 4 channels simultaneously inside the inner coefficient loop, so as to amortize by the same factor the cost of coefficient loads. Ablation studies reveal using wider float2 to float4 datatypes in the coefficient loop nearly double the throughput compared to the non-vectorized kernel, as can be expected from the shared memory traffic per operand coupling: (2 × 4 + 16) = 24 bytes loaded per batch, versus (4 × 2 × 4 + 16)/4 = 12 bytes loaded per batch with 4-fold vectorization, with 16B coefficients. • Coefficient packing: using narrow index data types so as to halve the size of each coefficient load when dimensions are small enough. The array of coefficients C is passed as a pointer to Coef structures holding one 32-bit float value and three indices of type Idx, where: – Idx = uint8 if all I/O dimensions are bounded by 255. Coef is then a 56 bit type, aligned to 64 bits, which can be loaded in a single LDS.64 instruction. – Idx = int32 by default. Coef is then a 128 bit type, which can be loaded in a single LDS.128 instruction. In practice, narrow indices have to be expanded to 32 bits in CUDA registers, and the upside of index narrowing is mostly that of a reducing shared memory pressure. • Coefficient distribution and occupancy: our kernels buffer input rows in shared memory to avoid long-scoreboard stalls when loading operands during the inner loop of algorithm 1. Since SMEM footprint limits the number of resident blocks on each SM, maximizing the number of resident threads requires increasing the block sizes above the given channel count. To that end, the coefficient loop is parallelized along the y-axis of a block, by first grouping coefficients by output indices before splitting them in evenly-sized groups.
Algorithm 1: Tensor product evaluation: parallel trailing channels, sequential aggregation Data: non-zero COO coefficients C = (i0, i1, i2, val) ∈ (N3 × R)Nc sorted by output index, ′ l.h.s. x ∈ Rd×Nq , r.h.s. y ∈ Cd ×Nq , thread index t ∈ N D Result: output z = x ⊗ y ∈ R i0 ← 0 ; zi0 ← 0; for p = 0 . . . Nc − 1 do coef ← LOAD C[p] ; i1, i2 ← coef.i1, coef.i2 ; xi1 ← LOAD x[i1, t] ; yi2 ← LOAD y[i2, t] ; if coef.i0 == i0 then zi0 + = coef.val ∗ xi1 ∗ yi2 ; end else z[i0, t] ← STORE zi0 ; zi0 ← coef.val ∗ xi1 ∗ yi2 ; i0 ← coef.i0 end end
Backward pass. Back-propagating through the tensor_product primitive can be implemented as two additional calls to the same tensor_product kernel: hence yielding 3n kernel calls for a model with n tensor products. This is made possible by the fact that algorithm 1 is generic with respect to the sparse coefficient array, see equation (15). Although this simple backward algorithm provides infinite differentiability, it has the disadvantage of loading output cotangents δz twice in on-device memory, which can be significantly expensive as the output feature dimensions scales as the product of input feature dimensions. Our backward tensor product kernel streams through batches of (x, y, dz) once and computes the two cotangents as (15), by reusing the same device code for the bilinear coupling (algorithm 1). 25
Preprint
forward
forward
forward
backward
Figure 11: Computation graphs for the tensor product forward pass (a) and backward pass (b-c). Back-propagating through a tensor product operation yields two transposed tensor product operations, by the Leibniz rule. Instead of calling two tensor product kernels, a dedicated backward kernel reuses device code for batch-wise bilinear coupling, while streaming through batches of the dz cotangents only once and thus saving memory traffic.
E.2.2
CUDA M ESSAGE PASSING C ONVOLUTION
It is worth noting that our CUDA message passing kernel is not constrained to edge based spherical harmonic embedding as second operand. Many Euclid equivariant MLIP architectures use this message-passing update of node features, such as MACE [3] and NequIP [2], and message formation is a known bottleneck of most MLIP architectures.8 For practical improvements, it may however be more important to consider hardware behavior and constraints at typical values of lmax rather than complexity arguments on hyperparameters. Even with an ideal memory-bound tensor product kernel, the message-passing operation may imply materializing edge features m̃ab = xa ⊗ Yl (rab ) in global memory. Their leading axis, the number of edges, is more than one order of magnitude larger than the number of atoms N (average number of neighbors of around 45 in organic molecules at 5Å cutoff). Our CUDA convolution kernel streams through edges and computes messages via a straightforward i0 overload of algorithm 1 to three operands given sparse 4D coefficients Ci1i2i3 , which we refer to as the bigotimes device routine. See equation (13) and figure 12. The trilinear mixing is embedded within an outer loop over receiver nodes, and an inner loop over sender nodes (neighbors) which relies on a Compressed Sparse Row (CSR) representation of the adjacency matrix. During each neighbor loop, messages m̃ab are accumulated in a shared memory buffer. While this accumulation strategy and CSR adjacency format enables efficient and deterministic message aggregation (free of compare-and-swap operations, a.k.a. atomics), the additional operand and message buffers increase the shared memory pressure compared to the raw tensor product kernel.
Backward pass. The current backward kernel processes each of the trilinear mixing operations of Eq. (16) with a common bigotimes() routine (see Fig. 12). While postponing the scalar mixing step could save a few FMUL/FMA operations in the forward pass, treating scalar mixing as a distinct operation in the backward pass leads to more complex device code and accumulation patterns (inner product reduction on channels of dyab ) and increased register pressure. While backward convolution still shows a good margin for optimization, relative algorithmic simplicity remains a constraint for high enough occupancy on the device.
8 MACE can be a notable exception due to its expensive multi-body atomic cluster expansion. However, with correlation 3 and maximal degree 1 for output node features, about 80% of runtime can be spent on the message-passing operation.
26
Preprint
Figure 12: Equivalent computation graphs for the convolution operation: the edge-wise composition of a scalar mixing on top of a bilinear tensor product (left) is equivalent to an edge-wise trilinear coupling (right) by associativity. E.3
PALLAS GPU CONVOLUTION KERNEL
The Pallas GPU kernel implements the same message-passing operation as Eq. 12, but differs from the CUDA implementation in how the sparse Clebsch–Gordan contraction and edge traversal are scheduled. Its main distinguishing feature is trace-time specialization of the Clebsch–Gordan coefficients. Edges are additionally traversed in receiver order, which trades receiver-local, atomicsfree forward aggregation against locality and cache reuse in computations where atomic accumulation is required. Algorithm 2 summarizes the forward kernel. Note that buffering of xa , yab and sab through SMEM is kept implicit for conciseness, and the LOAD directive refers to a shared memory load (LDS) as in algorithm 1. Trace-time Clebsch–Gordan specialization. Unlike the AOT-compiled CUDA Tensor Product implementation of Algorithm 1, the Pallas kernel does not loop through an array of coefficients at runtime. The non-zero coefficients and their indices depend only on the selected I/O feature spaces, and are therefore statically available for the Mosaic compiler9 . Writing the sparse records as (i0p , i1p , i2p , cp ), the loop over p is unrolled at trace time and the corresponding contraction paths are embedded directly in the generated program. There is therefore no LOAD of coefficients and their indices at runtime, in contrast with Algorithm 1, and the operand values can be loaded at each iteration without waiting for the coefficient and its indices to become available in registers. The records are grouped by output index i0, and passed as a Python dictionary – which forces inlining by the upstream HLO compiler – mapping i0 to a sequence C[i0] of triplets (i1p , i2p , cp ) in Algorithm 2. For a fixed edge, output feature i0, and channel q, the contraction X i0 i2 Ci1,i2 xi1,q yab a i1,i2 i3(i0),q
can consequently be accumulated in registers before applying the corresponding edge scalar sab . Trace-time specialization removes the coefficient loads and run-time loop and indexing overhead associated with traversing the sparse representation. This leads to having mostly FMA and LDS instructions that are all visible to the compiler, and can be reordered opportunistically to hide the latency of these instructions. Receiver-ordered aggregation. Edges are stored in compressed sparse row (CSR) format ordered by receiver b, such that [rowptr[b], rowptr[b + 1]) contains all edges incident on b. In the forward 9 Clebsch-Gordan coefficients grouped by indices are actually passed as nested Python dictionaries in the kernel source, which forces inlining by the upstream HLO compiler.
27
Preprint
Algorithm 2: Pallas GPU message-passing convolution: receiver-ordered aggregation and trace-time Clebsch–Gordan specialization. Data: static coefficients grouped by output index C[i0] = (i1p , i2p , cp )p∈Pi0 ; l.h.s. node ′ ′′ features x ∈ RN ×d×Nq ; r.h.s. edge features y ∈ RE×d ; edge scalars s ∈ RE×d ×Nq ; sender indices sender; receiver CSR pointers rowptr; channel q Result: aggregated messages m ∈ RN ×D×Nc foreach receiver b do ebegin , eend ← rowptr[b], rowptr[b + 1] acc ← 0 for e = ebegin , . . . , eend − 1 do a ← sender[e] // UNROLL for i0 = 0, . . . , D − 1 do z←0 for (i1p , i2p , cp ) ∈ C[i0] do x ← LOAD xa [i1p , q] y ← LOAD yab [i2p ] z += cp ∗ x ∗ y end s ← LOAD sab [i3(i0), q] acc[i0, q] += s ∗ z end end mb ← STORE acc end return m
pass, one cooperative thread array (CTA) processes the complete reduction for a receiver. Its accumulator is initialized in registers, updated while traversing the incoming edges, and written to GMEM after the final edge. This avoids inter-CTA atomic aggregation of mb . In an edge-parallel implementation, several CTAs may contribute simultaneously to mb += m̃ab , requiring atomic read–modify–write operations. Receiver-stationary aggregation instead keeps the partial sum private to one CTA, providing a fixed receiver-local accumulation order and avoiding repeated stores of partial results. The same property does not hold generally in the backward kernel. Receiver ordering makes consecutive edges reuse the receiver cotangent dmb , improving temporal locality and potentially L2-cache reuse, but gradient contributions need not remain local to one receiver. The implementation therefore uses atomic accumulation where multiple CTAs contribute to the same dx, dy, or ds output. Receiver ordering thus reflects a trade-off between atomics-free, deterministic receiver aggregation in the forward pass and improved locality in computations whose output aggregation may require atomics. Edge traversal. The receiver CSR interval is traversed sequentially, or in small contiguous tiles. Pallas’s emit_pipeline can stage receiver-ordered edge-local operands from GMEM to SMEM using asynchronous transfers, including TMA, while the current tile is processed. The sender access xa remains an indirect gather. This pipelining is a scheduling optimization rather than a defining feature of the kernel: its benefit depends on the balance between memory traffic and tensor-product computation. The central consequences of receiver ordering are instead the forward aggregation strategy and the locality obtained when repeatedly accessing receiver-associated quantities. 28
Preprint
Algorithm 3: Pallas GPU message-passing convolution, backward pass: receiver-ordered traversal and trace-time Clebsch-Gordan specialization. Data: static coefficients grouped by output index C[i0] = (i1p , i2p , cp )p∈Pi0 ; l.h.s. node ′ ′′ features x ∈ RN ×d×Nq ; r.h.s. edge features y ∈ RE×d ; edge scalars s ∈ RE×d ×Nq ; output cotangent dm ∈ RN ×D×Nq ; sender indices sender; receiver CSR pointers rowptr; channel q ′ ′′ Result: input cotangents dx ∈ RN ×d×Nq , dy ∈ RE×d , ds ∈ RE×d ×Nq dx ← 0 foreach receiver b do ebegin , eend ← rowptr[b], rowptr[b + 1] for e = ebegin , . . . , eend − 1 do a ← sender[e] dx, dy, ds ← 0 // UNROLL for i0 = 0, . . . , D − 1 do s ← LOAD sab [i3(i0), q] dm ← LOAD dmb [i0, q] dz ← s ∗ dm, z ← 0 for (i1p , i2p , cp ) ∈ C[i0] do x ← LOAD xa [i1p , q] y ← LOAD yab [i2p ] z += cp ∗ x ∗ y dx[i1p ] += cp ∗ y ∗ dz dy[i2p ] += cp ∗ x ∗ dz end ds[i3(i0)] += dm ∗ z end dsab [:, q] ← STORE P ds dyab ← STORE q dy // reduction over channels dxa ← ATOMIC_ADD dx end end return dx, dy, ds
E.4
PALLAS TPU KERNEL
Notation
The following notation is specific to the TPU implementation:
• B: number of edges in an edge tile • t: edge-tile index • at , bt : sender and receiver index vectors for tile t • yt , st : edge-feature and edge-scalar tiles • x̂t = x[at ]: sender features gathered for tile t • k = i3(i0): scalar-mixing block of output coordinate i0, which selects the edge scalar s[k] • acc, current and stage: state used by the segmented receiver or sender reduction. The following acronyms can also be useful, though for more details we refer readers to Appendix F: HBM denotes the TPU’s high-bandwidth memory, VMEM its vector memory, and SMEM its scalar memory. Tiling over edges rather than receivers. In contrast to the receiver-stationary Pallas GPU kernel, the TPU kernel tiles receiver-ordered edges into blocks of B. For each tile, VMEM holds the edge-local operands and messages, while the TensorCore’s SMEM holds the corresponding sender and receiver indices. The kernel copies yt and st from HBM to VMEM, loads at and bt into SMEM, 29
Preprint
gathers the sender features x̂t = x[at ] into VMEM, and computes the edge-levelPweighted messages m̃ab of Eq. 12. These messages are then reduced over receivers to form mb = a∼b m̃ab . The reduction state is retained across consecutive tiles, so a receiver spanning a tile boundary is accumulated before being written to HBM (Algorithm 4). Like the CUDA and Pallas GPU convolution kernels, this fuses message formation with aggregation and therefore avoids materializing the complete edge-leading message tensor in HBM. The TPU-specific distinction is the edge-tiled schedule: VMEM holds a block of edge messages, while SMEM carries the edge indices used for the gathers and segmented receiver reduction. Algorithm 4: Fused message-passing convolution on TPU. acc ← 0, current ← −1 for each tile t of B receiver-ordered edges do yt , st ← HBM to VMEM, at , bt ← HBM to SMEM x̂t ← HBM to VMEM x[at [k]], k = 0, . . . , B − 1 m̃t ← M ESSAGE T ILE(x̂t , yt , st ) R EDUCE F LUSH(bt , m̃t , acc, current) end F LUSH(current)
Coefficient packing and Clebsch–Gordan contraction. M ESSAGE T ILE evaluates the contraction path by path. The static records (i0p , i1p , i2p , cp ) are grouped on the host first by output coordinate i0p and then by l.h.s. input coordinate i1p . This allows each B × Nq tile x̂i1p to be loaded once and reused across the corresponding i2p paths (Algorithm 5). As in the Pallas GPU kernel, the sparse coefficient loops are unrolled at trace time and cp is embedded directly in the generated computation. Unlike Algorithm 1 for CUDA, the TPU kernel therefore does not load and interpret coefficient records at run time. i3(i0),q
Factoring out the edge scalar. The edge scalar sab in Eq. 13 is the same for every path of output coordinate i0, so the kernel pulls it out of the path sum, as in Eq. 12: X i3(i0),q i2,q i0 m̃i0,q Ci1,i2 xi1,q yab . a ab = sab i1,i2
|
{z
i0,q zab
}
It accumulates z first, then multiplies by the edge scalar once per output coordinate instead of once per path. Algorithm 5: M ESSAGE T ILE for i0 output coordinate do z←0 for i1 l.h.s. input coordinate of i0 do for (i2p , cp ) path of (i0, i1) do z += cp x̂[i1] y[i2p ] end end m̃[i0] ← z s[i3(i0)] end
// contraction before mixing
// scalar mixing
Reducing over receivers. R EDUCE F LUSH performs a segmented sum of the edge-message tile m̃t , of shape (D, B, Nq ). Because the edges are ordered by receiver, equal receiver indices form contiguous segments. The routine walks the B edges eight at a time, accumulates messages belonging to current, and flushes the completed receiver sum to mcurrent when the receiver changes (Algorithm 6). The variables acc and current persist across tiles. 30
Preprint
Algorithm 6: R EDUCE F LUSH Function ReduceFlush(bt , m̃, acc, current): for each chunk of 8 edges, receivers b[0, . . . , 7] do start ← 0 // entered when the chunk crosses a receiver boundary if b[0] ̸= current or b[0] ̸= b[7] then for j = 0, . . . , 7 with b[j] ̸= b[j − 1], b[−1] ≡ current do acc += m̃[ start ≤ edge < j ] // complete current receiver if current ≥ 0 then Flush(current) current ← b[j], start ← j end end acc += m̃[ edge ≥ start ] end Function Flush(b): P stage ← sublanes acc, acc ← 0 VMEM to HBM stage → mb
Backward. The backward follows the same edge-tiled organization as Algorithm 4. In addition to the edge-local yt and st , it gathers xa and the receiver cotangent dmb for each edge, replaces M ES SAGE T ILE by the sweep in Algorithm 7, and reduces the contributions to dxa using R EDUCE F LUSH keyed on senders. This segmented reduction therefore uses a sender-ordered edge traversal. The cotangents dyab and dsab remain edge-local. The sweep makes one pass over the same static coefficient records to compute the three cotangents of k,q i0,q Eq. 16. Since m̃i0,q ab = sab zab with k = i3(i0), the edge scalar and the contraction also separate in the backward: X i0,q i0,q i0,q wab = sk,q dsk,q zab dmi0,q ab dmb , ab = b . i0 : i3(i0)=k
The sweep computes w once per output coordinate and uses it in both the dx and dy updates. The edge-scalar cotangent instead needs z, which the sweep recomputes from the same paths. Because the records are grouped by block k, each ds[k] is accumulated in registers and written once. The r.h.s. yab is shared across channels, so its cotangent is summed over q at the end. Algorithm 7: Backward sweep of one edge tile. dx̂, dy ← 0 for k scalar-mixing block do σ←0 for i0 output coordinate of block k do w ← s[k] dm[i0] z←0 for (i1p , i2p , cp ) path of i0 do z += cp y[i2p ] x̂[i1p ] dx̂[i1p ] += cp y[i2p ] w dy[i2p ] += cp x̂[i1p ] w end σ += z dm[i0] end ds[k] ← σ end P dy ← q dy
// mix the cotangent once
// as in the forward
// one write per block
VMEM budget. Each buffer has at most three axes: the tile’s B edges, the Nq channels, and the feature components of a node, edge, or message value, whose count grows with lmax . The 31
Preprint
forward holds the gathered node features and message tile, one edge-scalar array per irreduciblerepresentation block, and the segmented sum accumulator together with the row used for flushing; sender and receiver indices are held in the TPU TensorCore’s SMEM. The backward additionally holds the three cotangents and two copies of x and dm for pipelining. Together these buffers occupy 4.3 MiB in the forward and 5.4 MiB in the backward at lmax = 3, Nq = 128, and fp32, and roughly half as much at lmax = 2 for the same block sizes. In addition to these buffers, the unrolled contraction of Algorithm 5 uses VMEM for compiler spills from vector registers, taking the peak to 13.1 and 10.1 MiB of the TPU TensorCore’s 16 MiB budget (Table 4). 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.
F
pass
B
buffers
peak VMEM
fwd bwd
128 32
4.3 5.4
13.1 (82%) 10.1 (63%)
R ELEVANT GPU AND TPU ARCHITECTURAL CONCEPTS
In order to provide a more comprehensive description of our kernels, we begin with an overview of the relevant hardware concepts for both GPU and TPU, as well as their differences. This section is not intended as an exhaustive or even fully accurate description of these respective devices, but instead as a structured glossary of concepts useful to understand later sections. F.1
GPU AND TPU ACCELERATORS COMPARISON
GPUs are massively parallel architectures exposing a single-instruction, multiple-threads (SIMT) programming model that gives fine-grained control over every level of parallelism via the CUDA C++ language extension [11]. A typical server-grade GPU embeds about a hundred streaming multi-processors (SMs) on the same die, each SM being capable of scheduling 2048 threads and carrying a low-latency 256 KB memory bank (L1 cache) alongside a 256 KB register file. In contrast, TPUs rely on the Mosaic compiler to lower abstract array operations onto dedicated processing units, including a systolic matrix multiply unit (MXU) and single-instruction, multipledata (SIMD) vector processing unit (VPU) [37]. A TPU tray typically consists of 4 TPU chips, each containing one or two cores, each backed by large (tens of MB) dedicated memory banks. All these architectural differences translate into significant algorithmic differences between the optimized kernels on each platform, on which more details can be found in the sections below. In particular, since end-to-end performance can depend strongly on reducing memory traffic and avoiding the materialization of intermediate results through operation fusion, different hardware architectures may opt for different fusion strategies and very different sizes for temporary buffers. The algorithmic development of e3j, as outlined in appendix E reflects these differentiations. Furthermore, significant engineering effort has been recently made to optimize or diversify the compiler stack. Just-in-time (JIT) compilation, from either CUDA or domain-specific languages embedded in Python such as Pallas, allows the production of specialized device assembly code from static problem parameters. It is now an ubiquitous alternative to traditional ahead-of-time (AOT) compilation of CUDA source into SASS or PTX. Further details on these distinctions can be found in appendix D.3. F.2
OVERVIEW OF RELEVANT GPU CONCEPTS
Computational units. A GPU exposes a hierarchical execution model that maps a large number of software threads onto parallel hardware resources. It is useful to distinguish the programming 32
Preprint
hierarchy: threads, warps and cooperative thread arrays (CTAs), from the underlying hardware hierarchy of streaming multiprocessors (SMs) and their execution units. • At the hardware level, the GPU is composed of many streaming multiprocessors (SMs). Each SM contains several execution pipelines for floating-point and integer arithmetic, load/store operations, specialized functions, and matrix operations such as those implemented by Tensor Cores. • At the software level, a thread is one logical execution instance of the kernel program. Each thread maintains its own execution state, including registers and, on modern NVIDIA GPUs, its own program counter. • Threads are organized in CTA, which are partitioned into groups of 32 threads called warps. An SM schedules and issues instructions at warp granularity: typically, one instruction is issued to the active threads of a warp, which apply that instruction to their respective data. This execution model is referred to as Single Instruction, Multiple Threads (SIMT). Threads may follow different control-flow paths, although divergence within a warp generally reduces execution efficiency. • A cooperative thread array (CTA), also called a CUDA thread block, groups threads that cooperate on the same unit of work. All threads of a CTA are scheduled on the same SM, where they may synchronise and exchange data through shared memory. The SM partitions the CTA into warps and interleaves the execution of its ready warps. Multiple CTAs may reside concurrently on one SM when register and shared memory capacity permit. Memory hierarchy. GPU performance is strongly influenced by where data reside in the memory hierarchy. Storage closer to the execution units provides high bandwidth and low latency but limited capacity, motivating kernels to move data through successively smaller on-chip memories and maximize reuse before returning to global memory. • Registers provide the storage closest to arithmetic execution. Registers hold thread-local values such as operands, intermediate results and accumulators, and are allocated from an SM-wide register file among the resident threads. On H100, each SM provides 65 536 32-bit registers, corresponding to 256 KiB of register-file capacity. Their limited capacity makes register usage an important constraint: using more registers per thread or CTA can reduce the number of CTAs that can reside concurrently on an SM. • Shared memory (SMEM) is an explicitly managed on-chip scratchpad shared by the threads of a CTA. It enables fast data reuse and communication between cooperating threads. The H100 provides a combined 256 KiB L1/texture-cache and shared-memory pool per SM, of which up to 228 KiB can be configured as shared memory. Although L1 cache and SMEM use the same underlying capacity, they have different programming semantics: L1 is hardware-managed cache, whereas placement and access in SMEM are controlled explicitly by the kernel. • The L2 cache is a larger cache shared across the GPU and forms the last on-chip caching level before device memory. H100 GPUs provide a 50 MB L2 cache, allowing frequently accessed data to be retained on-chip and reducing repeated accesses to the substantially larger off-chip memory. • Global memory (GMEM) is the GPU-wide memory address space used to store the large tensors processed by a kernel. Its backing storage is off-chip high-bandwidth DRAM (HBM3 on the H100 SXM5, for example) and consequently has much greater capacity but higher access cost than registers, SMEM or the on-chip caches. Efficient kernels therefore aim to minimize global-memory traffic and to reuse data after bringing it on-chip. • On Hopper GPUs, the Tensor Memory Accelerator (TMA) provides a specialized mechanism for asynchronously transferring multidimensional blocks of data between GMEM and SMEM. Once a transfer has been initiated, computation may proceed independently while TMA moves the next block of data. Software pipelines can consequently overlap global-memory traffic with arithmetic, as exploited by the Mosaic GPU kernel described below. 33
Preprint
Performance limits. For many ML kernels, particularly those with low arithmetic intensity, the rate at which data can be transferred between HBM and the GPU is the dominant performance bottleneck. In this memory-bound regime, the peak HBM bandwidth provides a useful “speed-of-light” estimate: the minimum execution time is bounded by the number of bytes that must be transferred to and from HBM divided by the maximum sustainable HBM bandwidth. Kernel efficiency can therefore be assessed by comparing its achieved HBM bandwidth with the hardware peak. Kernels with sufficiently high arithmetic intensity may instead become compute-bound, in which case peak arithmetic throughput rather than HBM bandwidth determines the relevant performance ceiling. F.3
OVERVIEW OF RELEVANT TPU CONCEPTS
Computational units. A TPU exposes a different execution model from the SIMT hierarchy of a GPU. Rather than scheduling many independent software threads, a TPU TensorCore is programmed approximately as a sequential machine operating on wide vector tiles. • A TPU v6e chip contains one TensorCore, comprising a scalar unit, a vector processing unit (VPU), and two matrix-multiply units (MXUs). The scalar unit handles control, indexing and scalar arithmetic, the VPU performs general vector operations, and the MXUs provide high-throughput matrix multiplication. • The natural unit of vector computation is a two-dimensional vector-register tile. For 32bit values on TPU v6e, registers are organized as 8 × 128 elements (sublanes × lanes). Vector instructions operate on these tiles collectively rather than on independently scheduled threads. • Array layout is therefore closely tied to this native tile shape. Operations are most efficient when trailing dimensions map cleanly onto the 8 × 128 register layout; poorly aligned or very small arrays may waste vector capacity or require additional rearrangement. • The two MXUs are specialized systolic units for dense matrix multiplication, while the VPU handles the more general elementwise, reduction and permutation operations used throughout Pallas kernels. Memory hierarchy. Large tensors reside in off-chip HBM and are staged through softwaremanaged on-chip memories before computation. • Vector registers (VREGs) hold operands and intermediate array-valued results consumed by the VPU and MXUs. Their limited capacity makes register pressure and spilling important considerations. • Vector memory (VMEM) is a large software-managed on-chip scratchpad for array data. TPU v6e provides approximately 128 MiB of VMEM per TensorCore, allowing substantially larger working sets to remain on-chip than in GPU shared memory. • Scalar memory (SMEM) is a separate on-chip memory for scalar data such as indices and control information. TPU v6e provides approximately 1 MiB of SMEM. This should not be confused with GPU “SMEM”, which denotes shared memory. • HBM provides the large off-chip storage for kernel inputs and outputs. TPU v6e provides 32 GB of HBM with approximately 1.64 TB/s peak bandwidth. Data are transferred between HBM and VMEM using asynchronous DMA engines, allowing future tiles to be prefetched while the current tile is being processed. Performance limits. As on GPUs, low-arithmetic-intensity TPU kernels may be bounded by HBM bandwidth, whereas compute-intensive kernels may instead be limited by VPU or MXU throughput. Efficient TPU kernels therefore aim to retain data in VMEM, respect the native 8 × 128 vector layout, and overlap HBM–VMEM transfers with computation. F.3.1
R ELEVANT DISTINCTION BETWEEN GPU AND TPU ARCHITECTURES
The GPU and TPU kernels are shaped by fundamentally different execution and memory models. On a GPU, computation is organized around CTAs composed of SIMT warps, with thread-local registers and a relatively small CTA-local shared-memory scratchpad. By contrast, a TPU TensorCore 34
Preprint
operates on wide 8 × 128 vector-register tiles and provides a much larger software-managed VMEM working memory. On H100, at most 228 KiB of shared memory is available per SM, whereas TPU v6e provides approximately 128 MiB of VMEM per TensorCore. These differences lead naturally to different kernel organizations: the GPU kernel keeps a receiver reduction local to one CTA and streams edge data through SMEM, whereas the TPU kernel operates on larger VMEM-resident windows using tile-wide vector operations.
G
R ECONSTRUCTIONS
G.1
R ECONSTRUCTION OF HARMONIC POLYNOMIALS .
The basis (Ylm )−l≤m≤l ∈ Yl of spherical harmonic polynomials of order l depends on the basis (x, y, z) of R3 by the spin equations demanding that Ylm is an eigenvector of z-axis (infinitesimal) rotations, σz · Ylm = imYlm (19) Noting that σz can be written as the infinitesimal rotation r∂ϕ of the longitude angle ϕ, equation (19) implies that Ylm varies as eimϕ along the longitude angle ϕ. Up to a choice of normalization factor cl , for m = ±l, this implies that Yl±l is the fastest oscillating degree l-monomial with respect to ϕ: Yl±l (x, y, z) = cl (x + iy)l = cl rl cos(θ)l e±ilϕ . In general, polynomials of lower absolute spin Ylm with |m| < l similarly take the form
(20)
Ylm (x, y, z) = rl Plm (cos θ, sin θ) eimϕ , (21) 0 and is an eigenvector of σz with eigenvalue im (in particular, Yl is always invariant with respect to z-axis rotations). They can be obtained from Yl±l by iterating the ladder operators σ± = σx ± iσy , which increase / decrease the magnetic quantum number m by 1. Note that our choice of generators is in agreement with physicists and chemists’ quantum state notation |lm⟩, while real harmonics Ylm are not eigenvalues of σz : only |m| is determined. In practice, the computation of polynomials Ylm is cached and staged out of XLA compilation. It returns an integer-valued array of exponents Y.exp :: [M, 3], and a sparse coefficient matrix Y.coef :: [N, M ] where N is the number of polynomials in the current basis, and M the number of distinct monomials required for their evaluation. G.2
R ECONSTRUCTION OF C LEBSCH G ORDAN COEFFICIENTS
LM LL The irreducible CG array Clm,l ′ m′ can be reconstructed from Cll,l′ l′ , expressing the top-spin eigenvectors with respect to the pure generators. The total z-spin M = m + m′ can indeed be lowered by the ladder operator σ− , acting as σ− ⊗ 1 + 1 ⊗ σ− on the tensor product space by the Leibniz rule, and any eigenvector |LM ⟩ for M < L can be obtained up to a normalization factor by iterating σ− from |LL⟩.
Therefore one only need to solve the eigenvalue equations ⃗σ 2 |LL⟩ = L(L + 1)|LL⟩ and σz |LL⟩ = L|LL⟩ to obtain the top-spin generators of each irreducible subspace of Yl ⊗ Yl′ , for each value of L ∈ {|l − l′ |, . . . , l + l′ }. The eigenvalue equations can be expressed and easily solved as a triangular system, with the algorithm documented below. By the second order Leibniz rule, the operator J+ J− acts on a dyadic tensor product as: J+ J− = (J+ J− ⊗ 1) + (1 ⊗ J+ J− ) + (J+ ⊗ J− ) + (J− ⊗ J+ ) 2
Since J given by
(22)
= J+ J− + Jz + Jz2 , the action of squared angular momentum on a pure state |lm, l′ m′ ⟩ is
J 2 |lm, l′ m′ ⟩ = l(l + 1) + l′ (l′ + 1) + 2mm′ · |lm, l′ m′ ⟩ + c− (m)c′+ (m′ ) · |l(m − 1), l′ (m′ + 1)⟩ + c+ (m)c′− (m′ ) · |l(m + 1), l′ (m′ − 1)⟩, 35
(23)
Preprint
where the term in curly brackets contains the independent action of J 2 on each operand, and twice the action of Jz ⊗ Jz . P Assuming V = Vm,m′ |lm, l′ m′ ⟩ is an eigenvector |LL⟩ for some L, by additivity of the z-spin M = L we can write more succinctly: X X V = Vm · |lm, l′ (L − m)⟩ = Vm · |m, (L − m)⟩ (24) m
m
so that the eigenvalue equation reads as follows: L(L + 1) · Vm = {l(l + 1) + l′ (l′ + 1) + 2mm′ } · Vm + c − (m + 1)c′ + (m′ − 1) · Vm+1 ′
(25)
′
+ c + (m − 1)c − (m + 1) · Vm−1 . Given m is constrained to [−L, L], coordinates Vm of the max-spin eigenvector |LL⟩ can easily be constructed by solving the above triangular system in dense or sparse format.
36