
JAX 性能剖析完全指南使用 jax.profiler 进行时间追踪与设备内存分析【免费下载链接】jaxComposable transformations of PythonNumPy 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 profilingstart_server、start_trace、stop_trace、trace、annotate_function、TraceAnnotation、StepTraceAnnotation、register_subprocess用于捕获程序执行的时间线可在 Perfetto 或 XProf/TensorBoard 中查看 CPU、GPU、TPU 上的活动。设备内存剖析Device memory profilingdevice_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_linkTrue时阻塞程序直到你打开 Perfetto 链接完成加载import jax with jax.profiler.trace(/tmp/jax-trace, create_perfetto_linkTrue): # 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。几个重要的行为约束源码中均有明确实现同时只能运行一个 tracestart_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 userhost如果使用 Google Cloudgcloud 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例如捕获 500mspython -m jax.collect_profile 9999 500该命令支持的参数来自 jax/collect_profile.py 的 argparse 定义参数含义默认值port要连接的剖析服务器端口必填duration_in_ms捕获时长毫秒必填--log_dirdirtrace 输出目录不指定则写入临时目录临时目录--no_perfetto_link禁用捕获后弹出 Perfetto 链接关闭默认弹出--host剖析服务器所在主机127.0.0.1默认的采集选项是host_tracer_level2、device_tracer_level1、python_tracer_level1见 jax/collect_profile.py也可以追加额外的--keyvalue选项覆盖。trace 输出在日志目录的plugins/profile/子目录下以*.xplane.pb形式存放随后会被转换为trace.json.gz供 Perfetto 上传或直接交给 TensorBoard 分析。stop_server()用于关闭剖析服务器。XProf / TensorBoard 剖析除了 PerfettoJAX 官方推荐的另一个时间剖析方案是 XProfOpenXLA 生态的剖析工具它同时支持 TensorBoard 插件和独立运行两种形态能够查看 GPU/TPU 上的详细活动。安装pip install xprof如果已安装 TensorBoardxprof包会自动安装 TensorBoard Profiler 插件。注意只安装一个版本的 TensorFlow/TensorBoard否则可能触发下文多个 TensorBoard 安装章节描述的Duplicate plugins错误。若需配合 nightly 版 TensorBoardpip install tb-nightly xprof-nightlyXProf 与 TensorBoard 配合XProf 是 TensorBoard 剖析与 trace 捕获功能背后的底层工具。只要安装了xprofTensorBoard 中就会出现 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 CTRLC 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 CTRLC 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_numglobal_step): train_step() global_step 1时间线上会出现train xx事件使用加速器时设备时间线上也会同步出现。源码中StepTraceAnnotation.__init__以_r1调用父类见 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.partialfrom functools import partial partial(jax.profiler.annotate_function, nameevent_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 的集成测试。配置 ProfileOptionsstart_trace和trace都接受可选的profiler_options参数类型为jax.profiler.ProfileOptions用于细粒度控制剖析行为。典型场景是关闭所有 Python 与 host 层 traceimport jax options jax.profiler.ProfileOptions() options.python_tracer_level 0 options.host_tracer_level 0 jax.profiler.start_trace(/tmp/profile-data, profiler_optionsoptions) # 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_levelhost 侧活动 trace 等级。0完全关闭 hostCPUtrace1仅 trace 用户主动插桩的 TraceMe 事件2包含等级 1外加高层程序执行细节如昂贵的 XLA 操作默认值3包含等级 2外加更冗长的底层执行细节如廉价的 XLA 操作。device_tracer_level是否启用设备 trace。0关闭设备 trace1启用设备 trace默认值。python_tracer_level是否启用 Python 函数调用 trace。0关闭 Python 函数调用 trace默认值1启用 Python trace。TPU 高级选项tpu_trace_modeTPU trace 模式取值包括TRACE_ONLY_HOST只 trace hostCPU侧活动不收集设备 traceTRACE_ONLY_XLA只 trace 设备上的 XLA 层操作TRACE_COMPUTEtrace 设备上的计算操作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_eventsCUPTI callback API 收集的最大事件数默认2*1024*1024gpu_max_activity_api_eventsCUPTI activity API 收集的最大事件数默认2*1024*1024gpu_max_annotation_strings可收集的最大注解字符串数默认1024*1024gpu_enable_nvtx_tracking在 CUPTI 中启用 NVTX 追踪默认Falsegpu_enable_cupti_activity_graph_trace为 CUDA graphs 启用 CUPTI activity graph 追踪默认Falsegpu_pm_sample_counters逗号分隔的 GPU 性能监控指标字符串如sm__cycles_active.avg.pct_of_peak_sustained_elapsed使用 CUPTI 的 PM sampling 特性收集默认关闭gpu_pm_sample_interval_usCUPTI PM sampling 的采样间隔微秒默认500gpu_pm_sample_buffer_size_per_gpu_mb每个设备用于 PM sampling 的系统内存缓冲MB默认 64MB最大支持 4GBgpu_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_RestrictProfilingToAdminUsers0 | sudo tee -a /etc/modprobe.d/nvidia-kernel-common.conf sudo update-initramfs -u sudo reboot now远程机器剖析被剖析程序运行在远程机器时可以在远程机器上启动 TensorBoard再用 SSH 本地端口转发访问 Web UI默认端口 6006ssh -L 6006:localhost:6006 remote server addressGoogle 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/pproflatest安装后位于$GOPATH/bin/pprofGOPATH默认为~/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(backendNone)支持指定后端名称如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(fmemory{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_serverjax.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 PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考