AI 中文总结
研究如何将傅里叶神经算子扩展到高分辨率空间网格,核心方法是引入分布式截断谱变换(DTST)并实现为DRIFT,主要贡献是在多GPU上实现显著加速,减少通信时间,提升计算效率。
AI 中文摘要
傅里叶神经算子(FNO)为偏微分方程学习解算子,在推理时比传统数值求解器快几个数量级,使其成为高分辨率计算物理的有吸引力的替代方法。将FNO扩展到高分辨率空间网格需要跨GPU分布数据,但每个谱层核心的分布式FFT需要多个密集的全对全集合来通信完整的空间张量,而大部分系数会立即被丢弃。我们引入了分布式截断谱变换(DTST),它颠倒了这个顺序。每个GPU通过部分DFT在本地仅计算谱卷积使用的一小部分频率模式,两个集合将结果与仅依赖于该模式数量而非空间分辨率的有效载荷相结合。DTST产生与带截断的标准分布式FFT相同的谱系数,同时提供空间数据并行和谱权重模型并行。我们展示了DRIFT,一种用于分布式傅里叶神经算子的DTST的GPU实现,使用可分离的逐维基矩阵和高效的GPU到GPU通信。在跨4至32个GPU、多达8个节点(每个节点4个GPU)的3D+时间FNO上,DRIFT实现了前向传播加速38至64倍,训练加速37倍,相比于分布式FNO基线,将通信时间从正向传播时间的97%减少到6%以下,并且在更高分辨率下加速效果不断增加。
英文摘要
Fourier Neural Operators (FNOs) learn solution operators for partial differential equations and offer orders of magnitude speedup over traditional numerical solvers at inference time, which makes them attractive surrogates for high-resolution computational physics. Scaling FNOs to high-resolution spatial grids requires distributing the data across GPUs, but the distributed FFT at the core of each spectral layer requires multiple dense all-to-all collectives that communicate the full spatial tensor, only for most coefficients to be discarded immediately. We introduce the Distributed Truncated Spectral Transform (DTST), which reverses this order. Each GPU computes only a small subset of frequency modes used by the spectral convolution locally via a partial DFT, and two collectives combine the results with a payload that depends only on this mode count, not the spatial resolution. DTST produces spectral coefficients identical to the standard distributed FFT with truncation, while providing both spatial data parallelism and spectral weight model parallelism. We present DRIFT, a GPU implementation of DTST for distributed Fourier Neural Operators, using separable per-dimension basis matrices and efficient GPU-to-GPU communication. On a 3D+time FNO across 4--32 GPUs, on up to 8 nodes (4 GPUs/node), DRIFT achieves a forward-pass speedup of 38--64$\times$ and a 37$\times$ training speedup over the distributed FNO baseline, reducing communication time from 97\% to under 6\% of the forward-pass time, with growing speedups at higher resolution.