High-level GPU code: a case study examining JAX and OpenMP.

Nestor Demeure, Theodore S. Kisner, Reijo Keskitalo, R. C. Thomas, Julian Borrill, W. Bhimji · 2023

In recent years, a new scientific software design pattern has emerged, pairing a Python interface with high-performance kernels in lower-level languages. The rise of general-purpose GPUs necessitates the rewriting of many such kernels, which poses challenges in GPU programming and ensures future portability and flexibility. This paper documents our experience and observations during the process of porting TOAST, a cosmology software framework designed to take full advantage of a supercomputer, to work with GPUs. This exploration led us to compare two different porting strategies: utilizing the JAX Python library and employing OpenMP Target Offload compiler directives. JAX allows kernel code to be written in pure Python, whereas OpenMP Target Offload is a directive-based strategy that integrates seamlessly with our existing OpenMP-accelerated C++ kernels. Both frameworks are high-level, abstracting system architecture details while aiming for straightforward, portable, yet performant GPU code. Through the porting of a dozen kernels, we delve into the analysis of development cost, performance, and the viability of employing either of these frameworks for complex numerical Python applications.

Read the paper · More papers on PaperTik