分析设备内存#
注意
2025 年 6 月更新:我们建议使用 XProf 分析来进行设备内存分析。生成分析报告后,打开 Tensorboard 分析器的 memory_viewer 选项卡,以查看更详细且易于理解的设备内存使用情况。
JAX 设备内存分析器使我们能够探索 JAX 程序如何以及为何使用 GPU 或 TPU 内存。例如,它可用于:
找出在特定时间点 GPU 内存中有哪些数组和可执行文件,或者
追踪内存泄漏。
安装#
JAX 设备内存分析器输出的结果可以使用 pprof (google/pprof) 进行解读。首先,请按照安装说明安装 pprof。在撰写本文时,安装 pprof 需要先安装 1.16+ 版本的 Go 和 Graphviz,然后运行:
go install github.com/google/pprof@latest
这会将 pprof 安装为 $GOPATH/bin/pprof,其中 GOPATH 默认为 ~/go。
注意
来自 google/pprof 的 pprof 版本与作为 gperftools 软件包一部分发布的同名旧工具不同。gperftools 版本的 pprof 无法与 JAX 配合使用。
了解 JAX 程序如何使用 GPU 或 TPU 内存#
设备内存分析器的常见用途是弄清楚为什么 JAX 程序占用了大量 GPU 或 TPU 内存,例如在尝试调试内存溢出(OOM)问题时。
要将设备内存分析数据捕获到磁盘,请使用 jax.profiler.save_device_memory_profile()。例如,考虑以下 Python 程序:
import jax
import jax.numpy as jnp
import jax.profiler
def func1(x):
return jnp.tile(x, 10) * 0.5
def func2(x):
y = func1(x)
return y, jnp.tile(x, 10) + 1
x = jax.random.normal(jax.random.key(42), (1000, 1000))
y, z = func2(x)
z.block_until_ready()
jax.profiler.save_device_memory_profile("memory.prof")
如果我们先运行上述程序,然后执行:
pprof --web memory.prof
pprof 会打开一个网页浏览器,其中包含以下调用图(callgraph)格式的设备内存分析可视化结果:
调用图是对每个活动缓冲区分配时 Python 堆栈的可视化。例如,在此特定案例中,可视化显示 func2 及其被调用函数负责分配了 76.30MB 内存,其中 38.15MB 是在从 func1 调用 func2 的过程中分配的。有关如何解读调用图可视化结果的更多信息,请参阅 pprof 文档。
使用 jax.jit() 编译的函数对设备内存分析器而言是“不透明”的。也就是说,在 jit 编译函数内部分配的任何内存都将归于整个函数。
在该示例中,调用 block_until_ready() 是为了确保在收集设备内存分析数据之前 func2 已经完成。有关更多详细信息,请参阅 异步分发。
调试内存泄漏#
我们还可以使用 JAX 设备内存分析器通过 pprof 可视化两个不同时间点获取的设备内存分析数据之间的变化来追踪内存泄漏。例如,考虑以下将 JAX 数组累积到不断增长的 Python 列表中的程序。
import jax
import jax.numpy as jnp
import jax.profiler
def afunction():
return jax.random.normal(jax.random.key(77), (1000000,))
z = afunction()
def anotherfunc():
arrays = []
for i in range(1, 10):
x = jax.random.normal(jax.random.key(42), (i, 10000))
arrays.append(x)
x.block_until_ready()
jax.profiler.save_device_memory_profile(f"memory{i}.prof")
anotherfunc()
如果我们仅仅在执行结束时可视化设备内存分析(memory9.prof),可能无法直观看出 anotherfunc 中循环的每次迭代都在累积更多的设备内存分配。
pprof --web memory9.prof
afunction 内部的大型但固定分配在分析中占据主导地位,但它不会随时间增长。
通过使用 pprof 的 --diff_base 功能来可视化循环迭代之间内存使用的变化,我们可以确定为什么程序的内存使用量会随时间增加:
pprof --web --diff_base memory1.prof memory9.prof
可视化结果显示,内存增长可归因于 anotherfunc 内部对 normal 的调用。