发表机构
INFN, Sezione di Firenze; INFN, Sezione di Pisa; INAF, Osservatorio di Astrofisica e Scienza dello Spazio; INFN, Sezione di Bologna(意大利国家核物理研究所佛罗伦萨分部; 意大利国家核物理研究所比萨分部; 意大利国家天体物理研究所空间科学与天体物理观测站; 意大利国家核物理研究所博洛尼亚分部)
机构由 AI 辅助整理,请以论文原文为准。AI 中文总结
本工作提出JAX生态中机器学习引力波代理模型mlgw和mlgw-bns的硬件加速实现,实现JIT编译、GPU加速和自动微分,在CPU和GPU上分别获得超过一个数量级和约两个数量级的加速,并成功集成到贝叶斯参数估计流程中。
AI 中文摘要
致密双星并合的精确且计算高效的波模型是引力波数据分析的基本要求。本工作提出了在JAX生态系统中机器学习代理模型mlgw和mlgw-bns的可微分且硬件加速的实现。所提出的框架支持JIT编译、向量并行化、原生GPU加速和自动微分,同时保留了原始代理训练基础设施及其对任意非进动波近似器的适用性。在二元黑洞情形下,使用新训练的SEOBNRv5HM近似器代理模型进行了基准测试;在CPU上,与原始mlgw实现相比,波形评估时间的加速超过一个数量级;此外,当利用GPU加速和大批量向量化时,这些加速达到约两个数量级。在双中子星情形下使用TEOBResumSPA代理模型进行的基准测试在GPU上取得了类似的增益。我们基于JAX的代理建模框架可以通过GWgpu-jax集成到贝叶斯参数估计流程中,GWgpu-jax是一个开源软件包,将兼容JAX的引力波波形生成器连接到JAX原生的采样算法。我们通过基于BlackJax-NS构建的嵌套采样流程展示了这一能力,使用mlgw和mlgw-bns对GW150914和GW170817事件进行了参数估计分析。这些分析在单个GPU上分别需要约十二分钟和十七分钟,并产生与LIGO-Virgo-KAGRA合作组织报告一致的后验分布。除了嵌套采样之外,JAX自动微分提供了波形模型的高效梯度和Hessian评估,从而实现了与基于梯度的贝叶斯推断方法的集成。
英文摘要
Accurate and computationally efficient waveform models of compact binary coalescences are a fundamental requirement for gravitational wave data analysis. This work presents differentiable and hardware-accelerated implementations of the machine-learning surrogate models mlgw and mlgw-bns within the JAX ecosystem. The proposed framework enables JIT compilation, vector paralellization, native GPU acceleration, and automatic differentiation, while preserving the original surrogate training infrastructure and its applicability to arbitrary non-precessing waveform approximants. Benchmarks are performed in the binary black hole case with a newly trained surrogate of the SEOBNRv5HM approximant; on CPU, they show speed-ups in waveform evaluation time with respect to the original mlgw implementation that exceed one order of magnitude; further, when exploiting GPU acceleration and large-batch vectorization these reach approximately two orders of magnitude. The benchmarks performed in the case of binary neutron stars with a TEOBResumSPA surrogate model achieve similar gains on GPU. Our JAX-based surrogate modeling framework can be integrated into Bayesian parameter estimation pipelines through GWgpu-jax, an open-source package that connects JAX-compatible gravitational waveform generators to JAX-native sampling algorithms. We demonstrate this capability with a nested-sampling pipeline built upon BlackJax-NS, performing parameter estimation analyses of the GW150914 and GW170817 events with mlgw and mlgw-bns. The analyses require approximately twelve and seventeen minutes, respectively, on a single GPU and yield posterior distributions consistent with those reported by the LIGO-Virgo-KAGRA Collaboration. Beyond nested sampling, JAX automatic differentiation provides efficient gradient and Hessian evaluations of the waveform models, enabling integration with gradient-based Bayesian inference methods.
Comments21 pages, 10 figures