如何理解与调试JAX内存占用?GPU内存溢出问题排查
JAX GPU内存暴增问题排查与解决
我是JAX新手,用它在GPU上运行X射线衍射数据索引的点云规则网格搜索代码。当test_mats维度设为[400000,3,3]时内存占用约15MB,但改成[500000,3,3]后程序提示需分配19GB内存,直接触发内存不足报错。怀疑是jit/vmap函数内部生成了(N, 3, gvec.shape[1])的临时数组,但无法查看函数内部内存情况。
示例代码
import sys import os import jax import jax.random import jax.profiler print('jax.version.__version__',jax.version.__version__) import scipy.spatial.transform import numpy as np # (3,N) integer grid spot positions hkls = np.mgrid[-3:4, -3:4, -3:4].reshape(3,-1) Umat = scipy.spatial.transform.Rotation.random( 10, random_state=42 ).as_matrix() a0 = 10.13 gvec = np.swapaxes( Umat.dot(hkls)/a0, 0, 1 ).reshape(3,-1) def count_indexed_peaks_hkl( ubi, gve, tol ): """ See how many gve this ubi can account for """ hkl_real = ubi.dot( gve ) hkl_int = jax.numpy.round( hkl_real ) drlv2 = ((hkl_real - hkl_int)**2).sum(axis=0) npks = jax.numpy.where( drlv2 < tol*tol, 1, 0 ).sum() return npks def testsize( N ): print("Testing size",N) jfunc = jax.vmap( jax.jit(count_indexed_peaks_hkl), in_axes=(0,None,None)) key = jax.random.PRNGKey(0) test_mats = jax.random.orthogonal(key, 3, (N,) )*a0 dev_gvec = jax.device_put( gvec ) scores = jfunc( test_mats, gvec, 0.01 ) jax.profiler.save_device_memory_profile(f"memory_{N}.prof") os.system(f"~/go/bin/pprof -top {sys.executable} memory_{N}.prof") testsize(400000) testsize(500000)
运行输出
gpu4-03:~/Notebooks/JAXFits % python mem.py jax.version.__version__ 0.4.16 Testing size 400000 File: python Type: space Showing nodes accounting for 15.26MB, 99.44% of 15.35MB total Dropped 25 nodes (cum <= 0.08MB) flat flat% sum% cum cum% 15.26MB 99.44% 99.44% 15.26MB 99.44% __call__ 0 0% 99.44% 15.35MB 100% [python] 0 0% 99.44% 1.53MB 10.00% _pjit_batcher 0 0% 99.44% 15.30MB 99.70% _pjit_call_impl 0 0% 99.44% 15.30MB 99.70% _pjit_call_impl_python 0 0% 99.44% 15.30MB 99.70% _python_pjit_helper 0 0% 99.44% 15.35MB 100% bind 0 0% 99.44% 15.35MB 100% bind_with_trace 0 0% 99.44% 15.30MB 99.70% cache_miss 0 0% 99.44% 15.30MB 99.70% call_impl_cache_miss 0 0% 99.44% 1.53MB 10.00% call_wrapped 0 0% 99.44% 13.74MB 89.51% deferring_binary_op 0 0% 99.44% 15.35MB 100% process_primitive 0 0% 99.44% 15.30MB 99.70% reraise_with_filtered_traceback 0 0% 99.44% 15.35MB 100% testsize 0 0% 99.44% 1.53MB 10.00% vmap_f 0 0% 99.44% 15.31MB 99.74% wrapper Testing size 500000 2023-12-14 10:26:23.630474: W external/tsl/tsl/framework/bfc_allocator.cc:296] Allocator (GPU_0_bfc) ran out of memory trying to allocate 19.18GiB with freed_by_count=0. The caller indicates that this is not a failure, but this may mean that there could be performance gains if more memory were available. Traceback (most recent call last): File "~/Notebooks/JAXFits/mem.py", line 38, in <module> testsize(500000) File "~/Notebooks/JAXFits/mem.py", line 33, in testsize scores = jfunc( test_mats, gvec, 0.01 ) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ jaxlib.xla_extension.XlaRuntimeError: RESOURCE_EXHAUSTED: Out of memory while trying to allocate 20596777216 bytes. -------------------- For simplicity, JAX has removed its internal frames from the traceback of the following exception. Set JAX_TRACEBACK_FILTERING=off to include these.
问题根源
- jit与vmap顺序错误:原代码先对单样本函数jit,再套vmap,JAX会为整个批量生成计算图,导致中间临时数组(如
hkl_real)以完整批量尺寸存储。当N=50万时,hkl_real尺寸为(500000, 3, 343),若用默认float64存储,单个数组就占约3.8GB,叠加其他中间数组后内存需求暴增。 - 默认精度过高:JAX默认使用float64,内存占用是float32的两倍,而X射线衍射计算无需如此高的精度。
修复方案
1. 调整jit与vmap顺序
将vmap放在内部,jit放在外部,让JAX生成单个样本的计算图后批量复用,避免一次性生成大尺寸中间数组:
jfunc = jax.jit(jax.vmap(count_indexed_peaks_hkl, in_axes=(0, None, None)))
2. 改用float32降低内存占用
将输入数组转换为float32,JAX计算时会自动使用float32处理:
gvec = np.swapaxes(Umat.dot(hkls)/a0, 0, 1).reshape(3, -1).astype(np.float32) test_mats = test_mats.astype(jax.numpy.float32)
3. 简化不必要代码
去掉dev_gvec = jax.device_put(gvec),JAX会自动管理设备数据传输。
修复后完整代码
import sys import os import jax import jax.random import jax.profiler print('jax.version.__version__', jax.version.__version__) import scipy.spatial.transform import numpy as np # (3,N) integer grid spot positions hkls = np.mgrid[-3:4, -3:4, -3:4].reshape(3, -1) Umat = scipy.spatial.transform.Rotation.random(10, random_state=42).as_matrix() a0 = 10.13 # 改用float32 gvec = np.swapaxes(Umat.dot(hkls)/a0, 0, 1).reshape(3, -1).astype(np.float32) def count_indexed_peaks_hkl(ubi, gve, tol): """ See how many gve this ubi can account for """ hkl_real = ubi.dot(gve) hkl_int = jax.numpy.round(hkl_real) drlv2 = ((hkl_real - hkl_int)**2).sum(axis=0) npks = jax.numpy.where(drlv2 < tol*tol, 1, 0).sum() return npks def testsize(N): print("Testing size", N) # 调整jit和vmap顺序:先vmap再jit jfunc = jax.jit(jax.vmap(count_indexed_peaks_hkl, in_axes=(0, None, None))) key = jax.random.PRNGKey(0) test_mats = jax.random.orthogonal(key, 3, (N,)) * a0 # 转为float32 test_mats = test_mats.astype(jax.numpy.float32) scores = jfunc(test_mats, gvec, 0.01) jax.profiler.save_device_memory_profile(f"memory_{N}.prof") os.system(f"~/go/bin/pprof -top {sys.executable} memory_{N}.prof") testsize(400000) testsize(500000)
验证效果
调整后,N=50万时的内存占用会和N=40万时处于同一量级,不会出现19GB的内存需求。
内容的提问来源于stack exchange,提问作者Jon
相关产品推荐
相关产品推荐

