arXivDaily arXiv每日学术速递 周一至周五更新
arXiv周末暂无论文更新,休息一下吧,周末愉快~~

E3J:面向GPU和TPU的欧几里得等变操作的高效开源后端

E3J: An Efficient and Open-Source Backend for Euclidean Equivariant Operations on GPU and TPU

Olivier Peltre, Armand Picard, Adrien Pichard, Miguel Bragança, Luca Giacomoni, Valentin Heyraud, Zachary Weller-Davies, Christoph Brunken, Jules Tilly

arXiv 2609.35099首次发表:更新:

发表机构

InstaDeep; Prima Mente(InstaDeep公司; Prima Mente)

机构由 AI 辅助整理,请以论文原文为准。

AI 中文总结

e3j是一个开源的高效欧几里得等变后端,通过优化CUDA/Pallas内核和算法改进,在GPU和TPU上实现高性能,显著加速MLIP等几何深度学习任务。

AI 中文摘要

我们提出了e3j,一个用于几何深度学习应用的快速欧几里得等变后端,具有针对GPU和TPU的JAX绑定。通过利用优化的CUDA和Pallas内核以及算法改进,该库在前向和后向路径上均实现了最先进的吞吐量和运行时间。在机器学习原子间势(MLIP)用例中,它优于已建立的后端,在使用MACE的水盒子NPT模拟中,相比cuEquivariance测量到高达34%的加速,同时保持完全开源。e3j在张量积操作上实现了超过H100最大内存带宽80%的效率,并且在许多情况下,前向消息传递卷积的吞吐量相比以前可用的后端提高了一倍以上。此外,随着专用Pallas TPU内核的发布,e3j开启了在TPU架构上进行大规模等变深度学习工作负载的可能性,而这在以前是很难实现的。我们的基准测试表明,e3j还实现了超过TPUv6e内存带宽80%的性能,比e3nn-jax高出一个数量级。该库可在GitHub和PyPI上获取,并以开源Apache 2.0许可证发布。

英文摘要

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.

Comments9 pages (36 total), 12 figures, 4 tables

论文原文

arXiv 摘要页 · PDF 原文 · HTML 原文

↑