GPU 内存分配

GPU 显存分配#

当运行第一个 JAX 操作时,JAX 将预分配总 GPU 显存的 75%。 预分配可以最大限度地减少分配开销和内存碎片,但有时会导致内存溢出(OOM)错误。如果您的 JAX 进程因 OOM 而失败,可以使用以下环境变量来覆盖默认行为:

XLA_PYTHON_CLIENT_PREALLOCATE=false

这将禁用预分配行为。JAX 将改为按需分配 GPU 显存,从而可能降低整体内存使用量。然而,这种行为更容易产生 GPU 显存碎片,这意味着使用大部分可用 GPU 显存的 JAX 程序在禁用预分配的情况下可能会出现 OOM。

XLA_PYTHON_CLIENT_MEM_FRACTION=.XX

如果启用了预分配,此设置将使 JAX 预分配总 GPU 显存的 XX%,而不是默认的 75%。降低预分配量可以解决 JAX 程序启动时发生的 OOM 问题。

XLA_PYTHON_CLIENT_ALLOCATOR=platform

这使得 JAX 可以根据需要精确分配所需的显存,并释放不再需要的内存(请注意,这是唯一会释放 GPU 显存而不是重复利用它的配置)。此模式速度非常慢,因此不建议常规使用,但对于以尽可能小的 GPU 显存占用运行或调试 OOM 故障可能很有用。

OOM 故障的常见原因#

同时运行多个 JAX 进程。

可以使用 XLA_PYTHON_CLIENT_MEM_FRACTION 为每个进程分配适当的内存量,或者设置 XLA_PYTHON_CLIENT_PREALLOCATE=false

同时运行 JAX 和 GPU 版 TensorFlow。

TensorFlow 默认也会进行预分配,因此这类似于同时运行多个 JAX 进程。

一种解决方案是仅使用 CPU 版 TensorFlow(例如,如果您只使用 TF 进行数据加载)。您可以使用命令 tf.config.experimental.set_visible_devices([], "GPU") 禁止 TensorFlow 使用 GPU。

或者,使用 XLA_PYTHON_CLIENT_MEM_FRACTIONXLA_PYTHON_CLIENT_PREALLOCATE。也有类似的选项来配置 TensorFlow 的 GPU 显存分配(TF1 中的 gpu_memory_fractionallow_growth,应在传递给 tf.Sessiontf.ConfigProto 中设置。有关 TF2,请参阅 使用 GPU:限制 GPU 内存增长)。

在显示器所用的 GPU 上运行 JAX。

请使用 XLA_PYTHON_CLIENT_MEM_FRACTIONXLA_PYTHON_CLIENT_PREALLOCATE

禁用重计算 (Rematerialization) HLO pass

有时禁用自动重计算 HLO pass 是有利的,可以避免编译器做出较差的重计算选择。可以通过分别设置 jax.config.update('jax_compiler_enable_remat_pass', True)jax.config.update('jax_compiler_enable_remat_pass', False) 来启用或禁用该 pass。启用或禁用自动重计算 pass 会在计算和内存之间产生不同的权衡。但请注意,该算法比较基础,通常您可以通过禁用自动重计算 pass 并使用 jax.remat API 手动进行重计算,从而获得更好的计算与内存权衡。

实验性功能#

此处的功能均为实验性质,使用时需谨慎。

XLA_PYTHON_CLIENT_ALLOCATOR=vmm

这使用了 CUDA 的虚拟内存管理 (VMM) 分配器 (cudaDeviceAddressVmmAllocator)。这是一个仅适用于 CUDA 的实验性分配器,提供细粒度的虚拟内存控制。它不会预分配内存;请使用 XLA_CLIENT_MEM_FRACTION 来控制所使用的 GPU 显存比例。

TF_GPU_ALLOCATOR=cuda_malloc_async

这用 cudaMallocAsync 替换了 XLA 自己的 BFC 内存分配器。这将取消大的固定预分配,并使用一个会动态增长的内存池。预期的好处是不再需要设置 XLA_PYTHON_CLIENT_MEM_FRACTION

风险包括:

  • 内存碎片模式不同,因此如果您接近极限,由于碎片导致的 OOM 情况会有所不同。

  • 分配时间不会全部在启动时支付,而是在需要增加内存池时产生。因此,您可能会在启动时感受到速度不稳定性,对于基准测试,忽略前几次迭代会变得更加重要。

这些风险可以通过预分配一大块内存来缓解,同时仍能获得拥有增长内存池的好处。这可以通过 TF_CUDA_MALLOC_ASYNC_SUPPORTED_PREALLOC=N 来实现。如果 N 为 -1,它将预分配与默认分配量相同的大小。否则,它就是您想要预分配的字节大小。