You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何理解与调试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.

问题根源

  1. jit与vmap顺序错误:原代码先对单样本函数jit,再套vmap,JAX会为整个批量生成计算图,导致中间临时数组(如hkl_real)以完整批量尺寸存储。当N=50万时,hkl_real尺寸为(500000, 3, 343),若用默认float64存储,单个数组就占约3.8GB,叠加其他中间数组后内存需求暴增。
  2. 默认精度过高: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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.03 23:41:01