JAX 性能剖析完全指南:使用 jax.profiler 进行时间追踪与设备内存分析
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
导读
jax.profiler是 JAX 官方提供的性能剖析(profiling)模块,用于回答两个核心问题:程序的时间花在了哪里(CPU / GPU / TPU 上的追踪与时间剖析),以及设备的显存/内存被谁占用了(设备内存剖析与泄漏定位)。本文以仓库中 jax.profiler.rst 的 API 索引为骨架,结合 profiling.md、device_memory_profiling.md 两份官方指南以及 jax/_src/profiler.py 源码实现,系统讲解从程序化捕获、手动捕获、XProf/TensorBoard 可视化到pprof内存调用图分析的全链路实践,读完即可在自己的 JAX 程序上复现。
jax.profiler 模块概览
jax.profiler作为公开 API 通过 jax/init.py 的from jax import profiler as profiler导出,因此可以直接以jax.profiler.xxx方式调用。按 jax.profiler.rst 的划分,模块能力分为两大块:
- 时间剖析(Tracing and time profiling):
start_server、start_trace、stop_trace、trace、annotate_function、TraceAnnotation、StepTraceAnnotation、register_subprocess,用于捕获程序执行的时间线,可在 Perfetto 或 XProf/TensorBoard 中查看 CPU、GPU、TPU 上的活动。 - 设备内存剖析(Device memory profiling):
device_memory_profile、save_device_memory_profile,用于生成 pprof 格式的设备内存快照,回答"哪些数组和可执行对象此刻占用了 GPU/TPU 内存、它们在哪里被分配"以及"内存为何持续增长"。
从源码结构看,所有 API 的底层实现都集中在 jax/_src/profiler.py:时间剖析通过 C++ 扩展jax._src.lib._profiler的ProfilerServer、ProfilerSession、TraceMe完成,设备内存剖析则通过后端客户端的heap_profile()方法采集(见下文)。
时间剖析:用 Perfetto 查看程序追踪
上下文管理器jax.profiler.trace
最简单的用法是把待剖析的代码放进jax.profiler.trace上下文管理器。程序结束时,trace 会被写入指定的日志目录,并(在create_perfetto_link=True时)阻塞程序,直到你打开 Perfetto 链接完成加载:
import jax with jax.profiler.trace("/tmp/jax-trace", create_perfetto_link=True): # Run the operations to be profiled key = jax.random.key(0) x = jax.random.normal(key, (5000, 5000)) y = x @ x y.block_until_ready()运行结束后程序会打印一个指向ui.perfetto.dev的链接,浏览器打开后 Perfetto UI 会加载 trace 文件并打开可视化时间线。该链接只在首次打开时有效,打开后会重定向到一个长期有效的新 URL;点击 Perfetto UI 的 "Share" 按钮可以生成可供他人共享的 permalink。
start_trace/stop_trace程序化捕获
trace上下文管理器本质上只是start_trace与stop_trace的组合。查看 jax/_src/profiler.py 的源码可以看到,trace先调用start_trace(log_dir, ...),在finally块中调用stop_trace()。因此你也可以手动管理捕获窗口:
import jax jax.profiler.start_trace("/tmp/profile-data") # Run the operations to be profiled key = jax.random.key(0) x = jax.random.normal(key, (5000, 5000)) y = x @ x y.block_until_ready() jax.profiler.stop_trace()注意block_until_ready()的调用:JAX 采用异步派发,计算会排队后立即返回,如果不阻塞等待设备执行完成,trace 将无法捕获到 on-device 的执行部分。关于异步派发的原理可参见 async_dispatch.rst。
几个重要的行为约束(源码中均有明确实现):
- 同时只能运行一个 trace:
start_trace在_profile_state.profile_session非空时会抛出RuntimeError("Profile has already been started. Only one profile may be run at a time.")。 start_trace会调用xla_bridge.get_backend()确保后端先完成初始化,否则在 Cloud TPU 上 libtpu 尚未初始化会导致 TPU tracer 初始化失败、TPU 操作无法进入 profile(见 jax/_src/profiler.py)。- 捕获开始时,JAX 会自动写入元数据:
jax_version、jaxlib_version,以及每个已注册后端平台的{platform}_version(见 jax/_src/profiler.py)。 stop_trace将结果导出到log_dir;若设置了create_perfetto_trace或create_perfetto_link,还会把 trace 转换为perfetto_trace.json.gz文件(并删除 Perfetto 不支持的metadata字段),必要时启动一个监听127.0.0.1:9001的临时 HTTP 服务器供ui.perfetto.dev拉取文件。
远程剖析(Remote profiling)
当被剖析的程序运行在远程主机(如托管 VM)上时,需要建立 SSH 隧道转发 9001 端口,Perfetto 链接才能工作:
ssh -L 9001:127.0.0.1:9001 <user>@<host>如果使用 Google Cloud:
gcloud compute ssh <machine-name> -- -L 9001:127.0.0.1:9001手动捕获:profiler server +collect_profile
除了程序化捕获,还可以先在脚本中启动一个剖析服务器,再随时手动触发指定时长的捕获。这在剖析长运行程序(例如训练循环中的某一段)时特别有用:
import jax.profiler jax.profiler.start_server(9999)然后通过命令行工具jax.collect_profile(源码见 jax/collect_profile.py)触发捕获:
python -m jax.collect_profile <port> <duration_in_ms>例如捕获 500ms:
python -m jax.collect_profile 9999 500该命令支持的参数(来自 jax/collect_profile.py 的 argparse 定义):
| 参数 | 含义 | 默认值 |
|---|---|---|
port | 要连接的剖析服务器端口 | 必填 |
duration_in_ms | 捕获时长(毫秒) | 必填 |
--log_dir=<dir> | trace 输出目录;不指定则写入临时目录 | 临时目录 |
--no_perfetto_link | 禁用捕获后弹出 Perfetto 链接 | 关闭(默认弹出) |
--host | 剖析服务器所在主机 | 127.0.0.1 |
默认的采集选项是host_tracer_level=2、device_tracer_level=1、python_tracer_level=1(见 jax/collect_profile.py),也可以追加额外的--key=value选项覆盖。trace 输出在日志目录的plugins/profile/子目录下,以*.xplane.pb形式存放,随后会被转换为trace.json.gz供 Perfetto 上传,或直接交给 TensorBoard 分析。stop_server()用于关闭剖析服务器。
XProf / TensorBoard 剖析
除了 Perfetto,JAX 官方推荐的另一个时间剖析方案是 XProf(OpenXLA 生态的剖析工具),它同时支持 TensorBoard 插件和独立运行两种形态,能够查看 GPU/TPU 上的详细活动。
安装
pip install xprof如果已安装 TensorBoard,xprof包会自动安装 TensorBoard Profiler 插件。注意只安装一个版本的 TensorFlow/TensorBoard,否则可能触发下文"多个 TensorBoard 安装"章节描述的Duplicate plugins错误。若需配合 nightly 版 TensorBoard:
pip install tb-nightly xprof-nightlyXProf 与 TensorBoard 配合
XProf 是 TensorBoard 剖析与 trace 捕获功能背后的底层工具。只要安装了xprof,TensorBoard 中就会出现 "Profile" 标签页,用法与独立运行 XProf 完全一致(需要指向同一个日志目录):
tensorboard --logdir=/tmp/profile-data输出类似:
[...] Serving TensorBoard on localhost; to expose to the network, use a proxy or pass --bind_all TensorBoard 2.19.0 at http://localhost:6006/ (Press CTRL+C to quit)程序化捕获 + 查看 trace
用jax.profiler.start_trace/jax.profiler.stop_trace(或trace上下文管理器)把 trace 写到目录后,就可以让 XProf 指向同一目录进行查看。独立运行 XProf 的方式:
xprof --port 8791 /tmp/profile-data输出类似:
Attempting to start XProf server: Log Directory: /tmp/profile-data Port: 8791 XProf at http://localhost:8791/ (Press CTRL+C to quit)在浏览器中打开输出的 URL,左侧 "Runs" 下拉框选择运行,然后在 "Tools" 下拉框选择trace_viewer,即可看到执行时间线;支持 WASD 键导航,点击/拖拽选择事件可查看细节。
通过 XProf 手动捕获 N 秒 trace
步骤如下:
启动 XProf 服务器(默认端口 8791,可用
--port修改):xprof --logdir /tmp/profile-data/在要剖析的 Python 程序开头加入:
import jax.profiler jax.profiler.start_server(9999)XProf 会连接该剖析服务器。剖析长程序时放在程序开头即可;剖析短程序(如微基准)时,可以在 IPython 中启动剖析服务器,再在下一步开始捕获后用
%run运行短程序,或在程序开头用time.sleep()留出开始捕获的时间。打开
http://localhost:8791/,点击左上角 "CAPTURE PROFILE" 按钮,在 "profile service URL" 中填入localhost:9999(即上一步剖析服务器的地址),输入要剖析的毫秒数并点击 "CAPTURE"。如果被剖析代码尚未运行,在捕获进行期间运行它。
捕获完成后 XProf 自动刷新;在左侧 "Tools" 下选择
trace_viewer查看时间线。
XProf 还提供多项分析工具:Framework Op Stats、Graph Viewer、HLO Op Stats、Memory Profile、Memory Viewer、HLO Op Profile、Roofline Model 等。
自定义 trace 事件:给时间线添加标注
默认情况下 trace viewer 里的事件大多是 JAX 内部的底层函数。jax.profiler提供了三个 API 用于注入自定义事件,让时间线更具可读性。
TraceAnnotation上下文管理器
包裹一段代码,生成一个覆盖该代码段执行时长的 trace 事件:
import jax.numpy as jnp import jax.profiler x = jnp.ones((1000, 1000)) with jax.profiler.TraceAnnotation("my_label"): result = jnp.dot(x, x.T).block_until_ready()捕获期间,时间线上会出现名为my_label的事件。
StepTraceAnnotation:标记训练步
StepTraceAnnotation是TraceAnnotation的子类,专门用于标记训练步。除时间线事件外,剖析器还会为每个 step 事件提供性能分析;传入step_num关键字参数可以带上全局步号:
while global_step < NUM_STEPS: with jax.profiler.StepTraceAnnotation("train", step_num=global_step): train_step() global_step += 1时间线上会出现train xx事件;使用加速器时设备时间线上也会同步出现。源码中StepTraceAnnotation.__init__以_r=1调用父类(见 jax/_src/profiler.py)。
annotate_function装饰器
装饰一个函数,使其每次执行都被标记为同名 trace 事件(名称默认取__qualname__或__name__):
@jax.profiler.annotate_function def f(x): return jnp.dot(x, x.T).block_until_ready() result = f(jnp.ones((1000, 1000)))需要自定义事件名或附加参数时,用functools.partial:
from functools import partial @partial(jax.profiler.annotate_function, name="event_name") def f(x): return jnp.dot(x, x.T).block_until_ready()从 jax/_src/profiler.py 的实现可见,装饰器本质是在函数调用外包一层TraceAnnotation(name, **decorator_kwargs),decorator_kwargs会作为附加参数透传给 trace 事件。
在 tests/profiler_test.py 中,testTraceAnnotation与testTraceFunction验证了TraceAnnotation、裸装饰器以及partial传名/传 kwarg 三种用法都能正确执行且不改变函数行为。
register_subprocess:剖析子进程
当工作负载分布在多个独立进程中(例如 PyGrain 等数据加载 worker 可能影响主进程性能)时,可以把子进程的剖析服务器注册到当前进程:主进程收集 profile 时会把请求传播给所有已注册子进程的剖析服务器,并聚合它们的响应。注册后返回一个取消注册函数:
unregister = jax.profiler.register_subprocess(pid, port)需要注意:目前只支持子进程的 CPU 剖析(见 jax/_src/profiler.py 的 docstring)。tests/profiler_test.py 中有跨进程注册并聚合 trace 的集成测试。
配置 ProfileOptions
start_trace和trace都接受可选的profiler_options参数,类型为jax.profiler.ProfileOptions,用于细粒度控制剖析行为。典型场景是关闭所有 Python 与 host 层 trace:
import jax options = jax.profiler.ProfileOptions() options.python_tracer_level = 0 options.host_tracer_level = 0 jax.profiler.start_trace("/tmp/profile-data", profiler_options=options) # Run the operations to be profiled key = jax.random.key(0) x = jax.random.normal(key, (5000, 5000)) y = x @ x y.block_until_ready() jax.profiler.stop_trace()通用选项
host_tracer_level:host 侧活动 trace 等级。
0:完全关闭 host(CPU)trace;1:仅 trace 用户主动插桩的 TraceMe 事件;2:包含等级 1,外加高层程序执行细节,如昂贵的 XLA 操作(默认值);3:包含等级 2,外加更冗长的底层执行细节,如廉价的 XLA 操作。
device_tracer_level:是否启用设备 trace。
0:关闭设备 trace;1:启用设备 trace(默认值)。
python_tracer_level:是否启用 Python 函数调用 trace。
0:关闭 Python 函数调用 trace(默认值);1:启用 Python trace。
TPU 高级选项
tpu_trace_mode:TPU trace 模式,取值包括:TRACE_ONLY_HOST:只 trace host(CPU)侧活动,不收集设备 trace;TRACE_ONLY_XLA:只 trace 设备上的 XLA 层操作;TRACE_COMPUTE:trace 设备上的计算操作;TRACE_COMPUTE_AND_SYNC:同时 trace 设备上的计算操作与同步事件。
未指定时默认为
TRACE_ONLY_XLA。tpu_num_sparse_cores_to_trace:要 trace 的 TPU sparse core 数量;tpu_num_sparse_core_tiles_to_trace:每个 sparse core 内要 trace 的 tile 数量;tpu_num_chips_to_profile_per_task:每个 task 要剖析的 TPU 芯片数量;tpu_perf_counters:是否收集性能计数器,默认为True。
GPU 高级选项
gpu_max_callback_api_events:CUPTI callback API 收集的最大事件数,默认2*1024*1024;gpu_max_activity_api_events:CUPTI activity API 收集的最大事件数,默认2*1024*1024;gpu_max_annotation_strings:可收集的最大注解字符串数,默认1024*1024;gpu_enable_nvtx_tracking:在 CUPTI 中启用 NVTX 追踪,默认False;gpu_enable_cupti_activity_graph_trace:为 CUDA graphs 启用 CUPTI activity graph 追踪,默认False;gpu_pm_sample_counters:逗号分隔的 GPU 性能监控指标字符串(如"sm__cycles_active.avg.pct_of_peak_sustained_elapsed"),使用 CUPTI 的 PM sampling 特性收集;默认关闭;gpu_pm_sample_interval_us:CUPTI PM sampling 的采样间隔(微秒),默认500;gpu_pm_sample_buffer_size_per_gpu_mb:每个设备用于 PM sampling 的系统内存缓冲(MB),默认 64MB,最大支持 4GB;gpu_num_chips_to_profile_per_task:每个 task 要剖析的 GPU 数量;未指定、为 0 或非法值时剖析全部可用 GPU,可用于减小 trace 体积;gpu_dump_graph_node_mapping:是否把 CUDA graph 节点映射信息写入 trace,默认False。
高级配置示例
高级选项通过advanced_configuration字典传入:
options = ProfileOptions() options.advanced_configuration = {"tpu_trace_mode": "TRACE_ONLY_HOST", "tpu_num_sparse_cores_to_trace": 2}若传入未识别的键或非法值,会返回InvalidArgumentError。
故障排查
GPU 剖析看不到设备 trace
运行在 GPU 上的程序,trace viewer 顶部应出现 GPU stream 的 trace。如果只看到 host trace,请检查日志中是否有以下错误。
Could not load dynamic library 'libcupti.so.10.1':把libcupti.so所在路径加入LD_LIBRARY_PATH(可用locate libcupti.so查找路径):
export LD_LIBRARY_PATH=/usr/local/cuda-10.1/extras/CUPTI/lib64/:$LD_LIBRARY_PATH设置后若仍报错,先检查 GPU trace 是否实际上已经出现在 trace viewer 中——该消息有时在一切正常时也会出现,因为它会在多个位置查找libcupti。
CUPTI_ERROR_INSUFFICIENT_PRIVILEGES:运行以下命令(需要重启):
echo 'options nvidia "NVreg_RestrictProfilingToAdminUsers=0"' | sudo tee -a /etc/modprobe.d/nvidia-kernel-common.conf sudo update-initramfs -u sudo reboot now远程机器剖析
被剖析程序运行在远程机器时,可以在远程机器上启动 TensorBoard,再用 SSH 本地端口转发访问 Web UI(默认端口 6006):
ssh -L 6006:localhost:6006 <remote server address>Google Cloud 环境下:
gcloud compute ssh <machine-name> -- -L 6006:localhost:6006多个 TensorBoard 安装
启动 TensorBoard 报ValueError: Duplicate plugins for name projector,通常是同时安装了多个 TensorFlow/TensorBoard 版本(tensorflow、tf-nightly、tensorboard、tb-nightly都自带 TensorBoard)。建议全部卸载后重装单一版本:
pip uninstall tensorflow tf-nightly tensorboard tb-nightly xprof xprof-nightly tensorboard-plugin-profile tbp-nightly pip install tensorboard xprof设备内存剖析:用 pprof 分析 GPU/TPU 内存
设备内存剖析用于探究 JAX 程序为何以及如何使用 GPU/TPU 内存,典型场景:确定某个时刻哪些数组和可执行对象驻留在 GPU 内存中、定位内存泄漏。device_memory_profile通过插桩 JAX 的设备端分配、为每次分配捕获 Python 栈来工作,插桩始终开启,API 只是负责抓取快照(见 jax/_src/profiler.py 的 docstring)。返回的是 gzip 压缩的 pprof 格式二进制协议缓冲区,可用 pprof 可视化。
安装 pprof
需要先安装 pprof:安装 Go 1.16+ 与 Graphviz 后执行:
go install github.com/google/pprof@latest安装后位于$GOPATH/bin/pprof(GOPATH默认为~/go)。注意:这里指的是 Google 的pprof,与gperftools包附带的同名旧工具不是同一个,后者无法配合 JAX 使用。
保存设备内存剖析
使用save_device_memory_profile(filename)把快照写入文件:
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 可视化:
pprof --http=: memory.prof浏览器中会出现以调用图(callgraph)形式呈现的设备内存剖析:
调用图是每个存活 buffer 分配时刻的 Python 栈可视化。例如上例中,func2及其被调函数负责分配了 76.30MB,其中 38.15MB 是在func1到func2的调用路径内分配的。两个值得注意的细节:
- 用
jax.jit编译的函数对设备内存剖析器是不透明的:jit函数内部分配的内存会整体归因于该函数。 block_until_ready()用于确保func2在收集剖析前已完成执行(异步派发机制参见 async_dispatch.rst)。
另外,device_memory_profile(backend=None)支持指定后端名称(如"gpu"、"tpu"),返回字节串而不是直接写文件;save_device_memory_profile只是它的便捷封装(见 jax/_src/profiler.py)。tests/profiler_test.py 的testDeviceMemoryProfile验证了其返回类型为bytes。
用 diff_base 调试内存泄漏
借助 pprof 的--diff_base特性对比两个时间点的剖析,可以定位随时间增长的内存。考虑一个把 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),增长原因并不明显:
pprof --http=: memory9.profafunction中那个大而固定的分配主导了整个剖析,但它在时间上并不增长。改用--diff_base对比循环早期与结束时的剖析:
pprof --http=: --diff_base memory1.prof memory9.prof可视化结果清晰地表明:内存增长归因于anotherfunc内部的normal调用——每个迭代都会在设备上累积新分配,从而定位到泄漏源头。
总结:jax.profiler 的两种工作模式
结合 jax.profiler.rst 与源码实现,可以把jax.profiler的能力归纳为两条主线:
- 时间剖析:用
trace/start_trace/stop_trace程序化捕获,或用start_server+jax.collect_profile/ XProf 手动捕获;用TraceAnnotation、StepTraceAnnotation、annotate_function为时间线添加语义标注;用register_subprocess聚合子进程剖析;用ProfileOptions精细控制 host/device/Python 三层 tracer 与 TPU/GPU 采集参数。结果在 Perfetto 或 XProf/TensorBoard 中查看。 - 设备内存剖析:
save_device_memory_profile输出 pprof 格式快照,配合 pprof 调用图与--diff_base对比,回答"谁占用了显存"与"内存为何增长"两个问题。
对于进一步深入,可以阅读 jax/_src/profiler.py 的完整实现(包括stop_and_get_fdo_profile与PGLEProfiler等面向 GPU FDO/PGLE 的高级能力)、命令行工具 jax/collect_profile.py 的参数解析,以及 tests/profiler_test.py 中的集成测试来验证各 API 的实际行为。
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考