JAX 性能剖析实战:从基准测试到计算追踪与设备内存分析

JAX 性能剖析实战:从基准测试到计算追踪与设备内存分析 JAX 性能剖析实战从基准测试到计算追踪与设备内存分析【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jaxJAX 程序写好后性能好不好、时间花在哪里、显存被谁占着是三个层层递进的问题。本文基于 JAX 官方 201 系列指南 docs/201/profiling.md 展开系统讲解 JAX 的基准测试benchmarking注意事项、基于 XProf 的计算追踪profiling computation、以及基于 pprof 的设备内存剖析profiling device memory三大主题。读完你将掌握如何正确地为 JAX 代码计时避开 JIT 编译、异步派发、32 位默认精度等陷阱、如何用jax.profiler采集并查看执行时间线、以及如何定位显存占用和内存泄漏。一、为什么 JAX 的计时和别的框架不一样JAX 代码的运行时行为由四个特性决定任何性能对比尤其是与 NumPy、PyTorch 等其他系统的对比都必须把它们考虑进去JAX 代码是 JIT 编译的。绝大多数 JAX 代码都可以写成支持 JIT 的形式编译后运行速度会快很多详见 docs/201/jit.md。要榨出最大性能应在最外层函数调用上应用jax.jit。需要特别注意的是第一次运行 JAX 代码必然更慢因为它正在被编译——即使你自己的代码里没有显式使用jitJAX 内置函数同样会被 JIT 编译首次调用依然有编译开销。JAX 采用异步派发asynchronous dispatch。调用一个 JAX 函数会立即返回一个未来值真正计算在设备上异步执行。因此计时时必须调用.block_until_ready()来确保计算确实已经完成详见 docs/201/jit.md 的 Asynchronous dispatch 一节。从源码看block_until_ready在 jax/_src/api.py 中实现它会遍历 pytree 的所有叶子对jax.Array逐个或多个批量调用设备端的阻塞等待。JAX 默认只使用 32 位 dtype。64 位计算比 32 位贵得多做性能对比时务必让双方精度一致否则比较没有意义控制 JAX 默认精度的方法见 docs/101/arrays.md。CPU 与加速器之间的数据传输耗时。如果你只想测量求值函数本身花了多久应当先把数据放到目标设备上而不是把数据传输时间也算进去。1.1 一个完整的 JAX vs NumPy 微基准把上述技巧组合起来就是一个正确的微基准。下面这个例子用 IPython 的%time/%timeit魔法分别测量 NumPy 与 JAX 的运行时间原文档示例import numpy as np import jax def f(x): # 被基准测试的函数NumPy 与 JAX 均可运行 return x.T (x - x.mean(axis0)) x_np np.ones((1000, 1000), dtypenp.float32) # 与 JAX 默认 dtype 一致 %timeit f(x_np) # 测量 NumPy 运行时间 # 测量 JAX 设备传输时间 %time x_jax jax.device_put(x_np).block_until_ready() f_jit jax.jit(f) %time f_jit(x_jax).block_until_ready() # 测量 JAX 编译时间首次调用 %timeit f_jit(x_jax).block_until_ready() # 测量 JAX 稳态运行时间这里的关键细节输入用dtypenp.float32与 JAX 默认精度对齐保证对比公平jax.device_put负责把数据搬到设备上.block_until_ready()确保传输完成后再开始计时对f_jit的首次调用测量的是编译时间%timeit多次重复测量才是稳态运行时间。基准测试回答总共花了多久但要回答时间花在了哪里就需要剖析profile。二、计算剖析用 XProf 查看执行时间线JAX 的主要剖析工具是XProf它可以记录并可视化程序执行的详细轨迹按设备划分的时间线、算子级统计、内存使用等。本节介绍 XProfTensorBoard 剖析、轻量的 Perfetto 集成以及 NVIDIA 的 Nsight 工具。2.1 安装 XProfXProf 既可以作为 TensorBoard 的插件使用也可以作为独立程序运行pip install xprof如果你已安装 TensorBoardxprof这个 pip 包会自动安装 TensorBoard Profiler 插件。注意只能安装一个版本的 TensorFlow 或 TensorBoard否则可能遇到下文提到的 duplicate plugins 错误。如果使用 nightly 版本的 TensorBoard 做剖析需要配套 nightly 版 XProfpip install tb-nightly xprof-nightly2.2 XProf 与 TensorBoardXProf 是 TensorBoard 中剖析与轨迹捕获功能背后的底层工具。只要安装了xprofTensorBoard 中就会出现 Profile 标签页。用它启动剖析与独立启动 XProf 完全等价只要指向同一个日志目录即可涵盖捕获、分析、查看全部功能。XProf 取代了此前推荐的tensorboard_plugin_profile插件。$ 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)2.3 程序化捕获start_trace / stop_trace 与 trace 上下文管理器可以用jax.profiler.start_trace和jax.profiler.stop_trace在代码里插桩把剖析轨迹写到指定目录——这个目录应当与启动 XProf 时使用的--logdir一致之后用 XProf 查看。import jax jax.profiler.start_trace(/tmp/profile-data) # 运行要被剖析的操作 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 是异步派发的必须调用它确保设备上的执行确实发生并被轨迹捕获原因见 docs/201/jit.md 的异步派发一节。也可以用上下文管理器jax.profiler.trace替代start_trace/stop_trace的成对调用import jax with jax.profiler.trace(/tmp/profile-data): key jax.random.key(0) x jax.random.normal(key, (5000, 5000)) y x x y.block_until_ready()从实现看这些 API 定义在 jax/_src/profiler.pystart_trace第 151 行会加锁防止并发采集、初始化后端、写入 JAX/jaxlib/各后端版本元数据并创建ProfilerSessiontrace第 307 行本质上是start_tracetry/finally中的stop_tracestop_trace第 271 行将轨迹导出到日志目录。值得留意的是start_trace还支持create_perfetto_link与create_perfetto_trace参数见下文 Perfetto 章节。2.4 查看轨迹捕获完成后可以直接用独立 XProf 命令启动剖析 UI指向日志目录$ 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如 http://localhost:8791/。可用的轨迹会出现在左侧的 Runs 下拉菜单中选中感兴趣的 run 后在 Tools 下拉菜单中选择trace_viewer就能看到执行时间线。可以用 WASD 键导航轨迹点击或拖拽选择事件查看细节。2.5 手动触发捕获从运行中的程序采集 N 秒轨迹适用于对正在运行的程序手动触发一次指定时长的捕获步骤如下启动 XProf 服务器xprof --logdir /tmp/profile-data/在浏览器打开 http://localhost:8791/可用--port指定其他端口。如果 JAX 程序运行在远程服务器上见下文远程机器上剖析。在被剖析的 Python 程序开头加入import jax.profiler jax.profiler.start_server(9999)这会启动一个供 XProf 连接的 profiler 服务器进入下一步之前必须保证它已运行。用完后可调用jax.profiler.stop_server()关闭。想剖析长程序的某段如训练循环把这段代码放在程序开头正常启动程序即可想剖析短程序如微基准可以在 IPython shell 中启动 profiler 服务器然后在下一步开始捕获后用%run运行短程序或者在程序开头启动服务器后用time.sleep()给自己留出启动捕获的时间。打开 http://localhost:8791/点击左上角 CAPTURE PROFILE 按钮在 profile service URL 处填入 localhost:9999即上一步启动的 profiler 服务器地址填入要剖析的毫秒数点击 CAPTURE。如果目标代码尚未运行例如服务器是在 Python shell 里启动的在捕获进行期间运行它。捕获结束后 XProf 会自动刷新。并非所有 XProf 的剖析功能都与 JAX 对接所以一开始可能看起来什么都没捕获到。在左侧 Tools 下选择trace_viewer即可看到执行时间线。除 trace viewer 外XProf 还提供以下工具Framework Op Stats框架算子统计Graph Viewer计算图查看器HLO Op StatsHLO 算子统计Memory Profile内存剖析Memory Viewer内存查看器HLO Op ProfileHLO 算子剖析Roofline ModelRoofline 模型2.6 添加自定义轨迹事件默认情况下 trace viewer 中的事件大多是 JAX 内部的底层函数。你可以在代码中用jax.profiler.TraceAnnotation和jax.profiler.annotate_function添加自己的事件from functools import partial import jax import jax.numpy as jnp # 上下文管理器方式标记一段代码 with jax.profiler.TraceAnnotation(my_label): result jnp.dot(x, x.T).block_until_ready() # 装饰器方式标记一个函数 jax.profiler.annotate_function def f(x): return jnp.dot(x, x.T).block_until_ready() # 也可以通过 partial 传参改名 partial(jax.profiler.annotate_function, nameevent_name) def g(x): return jnp.dot(x, x.T).block_until_ready()从源码看TraceAnnotationjax/_src/profiler.py是_profiler.TraceMe的封装事件跨度即上下文执行时长annotate_function第 389 行默认使用函数的__qualname__或__name__作为事件名。此外还有一个StepTraceAnnotation第 363 行用于标记训练步并支持传入step_num让剖析器按步给出性能分析。与之互补的工具是jax.named_scope它给其上下文内创建的操作附加名字。与 trace 注解不同这些名字会流入 jaxpr 和编译后的 HLO因此不仅出现在时间线上还会出现在 XProf 的算子级视图以及编译器 dump中。其实现位于 jax/_src/api.py本质是向 JAX 的 name stack 追加前缀既可用作上下文管理器也可用作装饰器jax.jit jax.named_scope(layer) def layer(w, x): with jax.named_scope(dot_product): logits w.dot(x) with jax.named_scope(activation): return jax.nn.relu(logits)2.7 配置剖析选项ProfileOptionsstart_trace接受可选的profiler_options参数用于精细控制剖析器行为其类型为jax.profiler.ProfileOptions定义于 jax/_src/profiler.py继承自底层_profiler.ProfileOptions。例如关闭所有 Python 与 host 轨迹import jax options jax.profiler.ProfileOptions() options.python_tracer_level 0 options.host_tracer_level 0 jax.profiler.start_trace(/tmp/profile-data, profiler_optionsoptions) # 运行要被剖析的操作 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 侧CPU活动的追踪级别。取值含义0完全关闭 hostCPU追踪1只追踪用户插桩的 TraceMe 事件2包含级别 1 加上高层程序执行细节如昂贵的 XLA 操作——默认值3包含级别 2 加上更冗长的底层执行细节如廉价的 XLA 操作device_tracer_level是否启用设备追踪。取值含义0关闭设备追踪1启用设备追踪——默认值python_tracer_level是否启用 Python 追踪。取值含义0关闭 Python 函数调用追踪——默认值1启用 Python 追踪高级配置TPU 选项tpu_trace_modeTPU 追踪模式。可选值TRACE_ONLY_HOST只追踪 host 侧CPU活动不采集设备TPU/GPU轨迹TRACE_ONLY_XLA只追踪设备上的 XLA 级操作TRACE_COMPUTE追踪设备上的计算操作TRACE_COMPUTE_AND_SYNC同时追踪设备上的计算操作与同步事件。未提供时默认TRACE_ONLY_XLA。tpu_num_sparse_cores_to_trace指定追踪的 TPU sparse core 数量。tpu_num_sparse_core_tiles_to_trace指定每个 sparse core 内要追踪的 tile 数量。tpu_num_chips_to_profile_per_task每个任务要剖析的 TPU 芯片数。tpu_perf_counters是否收集性能计数器默认True。更完整的高级剖析标志列表包括功耗监控和周期性计数器采样可参考 OpenXLA XProf 仓库中的 Advanced Profiler Options 文档。高级配置GPU 选项gpu_max_callback_api_eventsCUPTI callback API 收集的最大事件数默认2*1024*1024。gpu_max_activity_api_eventsCUPTI 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 Performance Monitoring 指标字符串用 CUPTI 的 PM sampling 采集例如sm__cycles_active.avg.pct_of_peak_sustained_elapsed。PM sampling 默认关闭。可用指标见 NVIDIA CUPTI 文档的 metrics 表。gpu_pm_sample_interval_usCUPTI PM sampling 的采样间隔微秒默认500。gpu_pm_sample_buffer_size_per_gpu_mb每个设备的系统内存缓冲区大小MB默认 64MB最大支持 4GB。gpu_num_chips_to_profile_per_task每个任务要剖析的 GPU 设备数。不指定、设为 0 或非法值时剖析所有可用 GPU。可用于减小轨迹采集体积。gpu_dump_graph_node_mapping启用后把 CUDA graph 节点映射信息 dump 进轨迹默认False。高级配置使用示例options ProfileOptions() options.advanced_configuration {tpu_trace_mode : TRACE_ONLY_HOST, tpu_num_sparse_cores_to_trace : 2}如果传入任何无法识别的键或非法选项值会返回InvalidArgumentError。2.8 常见问题排查TroubleshootingGPU 剖析问题在 GPU 上运行的程序trace viewer 顶部应该能看到 GPU stream 的轨迹。如果只看到 host 轨迹请检查程序日志/输出中是否有下列错误。错误一Could not load dynamic library libcupti.so.10.1完整报错形如dso_loader.cc:55] Could not load dynamic library libcupti.so.10.1; dlerror: ... cannot open shared object file以及cupti_interface_-Subscribe(...) failed with error CUPTI could not be loaded or symbol could not be found把libcupti.so所在路径加入环境变量LD_LIBRARY_PATH可用locate libcupti.so查找路径。例如export LD_LIBRARY_PATH/usr/local/cuda-10.1/extras/CUPTI/lib64/:$LD_LIBRARY_PATH如果设置后仍报这个错先看看 trace viewer 里 GPU 轨迹是否其实已经出现了——该消息有时在一切正常的情况下也会出现因为程序会在多个位置查找libcupti库。错误二failed with error CUPTI_ERROR_INSUFFICIENT_PRIVILEGES完整报错形如cupti_interface_-EnableCallback(...) failed with error CUPTI_ERROR_INSUFFICIENT_PRIVILEGES及ActivityDisable ... CUPTI_ERROR_NOT_INITIALIZED执行以下命令注意需要重启echo options nvidia NVreg_RestrictProfilingToAdminUsers0 | sudo tee -a /etc/modprobe.d/nvidia-kernel-common.conf sudo update-initramfs -u sudo reboot now远程机器上剖析如果被剖析的 JAX 程序运行在远程机器上一个方案是在远程机器上执行上述全部步骤尤其是把 TensorBoard 服务器也启动在远程机器上然后用 SSH 本地端口转发把 TensorBoard Web UI 从本地接到远程。转发默认 TensorBoard 端口 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通常是因为装了两个版本的 TensorBoard 和/或 TensorFlowtensorflow、tf-nightly、tensorboard、tb-nightly这几个 pip 包都包含 TensorBoard。卸载单个 pip 包可能导致tensorboard可执行文件被移除而难以恢复因此可能需要全部卸载后重装单一版本pip uninstall tensorflow tf-nightly tensorboard tb-nightly xprof xprof-nightly tensorboard-plugin-profile tbp-nightly pip install tensorboard xprof2.9 轨迹中该看什么常见特征信号以下是一些常见的轨迹特征及其通常的含义设备时间线上操作之间出现空隙设备空闲在等待 host。典型原因逐算子 eager 派发应把更多程序包进jax.jit、步间关键路径上的 Python 工作数据加载、日志、或把值取回 host 引发的同步包括调试打印见 docs/201/debugging.md。第一步特别长是编译每个 JAX 类型签名只发生一次属预期行为。如果反复出现说明你在反复 retracing见 docs/201/slow-compilation.md。一连串密集的小 kernel单算子开销占主导更大的 jitted 区域能让编译器做更多融合。长耗时的集合通信操作对分片程序而言all-gather、reduce-scatter、all-reduce等操作的时间是通信时间。把它与计算时间对比可判断程序是否通信受限若是则重新审视分片策略见 docs/201/sharding.md。内存压力用 Memory Profile 和 Memory Viewer 工具查看随时间变化的分配情况与峰值时的构成详见下文设备内存剖析。2.10 剖析分布式代码轨迹按进程捕获每个本地设备一条时间线因此单进程跑满主机全部 8 块 GPU或 TPU 主机的全部芯片时无需额外设置即可覆盖包括分片计算执行的集合操作。对于多进程程序需要在每个进程中启动 profiler 服务器XProf 的捕获对话框接受逗号分隔的 profiler-service 地址列表可以把多个 host 捕获进同一个 profile。多进程编程本身见系统文档 docs/501/multiprocess.md。三、轻量方案用 Perfetto 查看程序轨迹作为 XProf 的轻量替代JAX profiler 可以生成可在 Perfetto 可视化器中查看的轨迹无需安装任何东西即可快速交互查看。目前这种方式会阻塞程序直到链接被点击且 Perfetto UI 加载完轨迹。with jax.profiler.trace(/tmp/jax-trace, create_perfetto_linkTrue): # 运行要被剖析的操作 key jax.random.key(0) x jax.random.normal(key, (5000, 5000)) y x x y.block_until_ready()计算完成后程序会提示你打开一个指向ui.perfetto.dev的链接。打开后 Perfetto UI 会加载轨迹文件并打开可视化器。打开链接后程序继续执行。该链接打开一次后失效但它会重定向到一个长期有效的新 URL。你可以在 Perfetto UI 中点击 Share 按钮创建轨迹的永久链接与他人共享。从实现看jax/_src/profiler.pycreate_perfetto_link会先把轨迹转换为perfetto_trace.json.gz移除 Perfetto 不喜欢的metadata字段随后在127.0.0.1:9001上启动一个临时 HTTP 服务托管该文件程序阻塞直到 Perfetto UI 取走文件。3.1 远程剖析当剖析运行在远程如托管 VM的代码时需要为 9001 端口建立 SSH 隧道链接才能生效$ ssh -L 9001:127.0.0.1:9001 userhost或使用 Google Cloud$ gcloud compute ssh machine-name -- -L 9001:127.0.0.1:90013.2 手动捕获除了用jax.profiler.trace以编程方式捕获外也可以在目标脚本中调用jax.profiler.start_server(port)启动剖析服务器如果只需要服务器在脚本的某一段生效结束后调用jax.profiler.stop_server()即可。脚本运行且 profiler 服务器启动后手动捕获并追踪$ python -m jax.collect_profile port duration_in_ms参数说明对应 jax/collect_profile.py 的实现portprofiler 服务器端口必填位置参数duration_in_ms捕获时长毫秒必填位置参数--log_dirdirectory of choice轨迹输出目录。默认输出到临时目录--no_perfetto_link禁用打开ui.perfetto.dev链接的提示。默认情况下程序会提示你打开 Perfetto 链接--host要捕获轨迹的主机默认127.0.0.1。该脚本还支持透传各类剖析选项如--host_tracer_level2 --device_tracer_level1 --python_tracer_level1与ProfileOptions的默认值对应由_parse_xprof_flags解析并合并进 XProf 采集选项。另外也可以把 TensorBoard 指向log_dir来分析轨迹见上文 XProf 与 TensorBoard 一节。四、NsightGPU 专用剖析NVIDIA 的Nsight工具可用于在 GPU 上追踪和剖析 JAX 代码。NVIDIA 官方文档提供了详细说明。需要留意的是Nsight 与 JAX 的 CUPTI 剖析可能存在订阅冲突源码 jax/_src/profiler.py 的 PGLE 路径对此有专门警告两者并存时可能采集到空轨迹。五、设备内存剖析显存被谁占用了对于绝大多数设备内存问题尤其是程序为什么 OOM 了推荐工具是 XProf按上文的轨迹捕获方式采集 profile 后打开 Memory Profile 和 Memory Viewer 工具即可看到随时间变化的内存使用、峰值时分配的构成以及每个 buffer 的大小和生命周期。JAX 还有一个互补的内存剖析器视角不同它给出设备上每个存活 buffer 的快照并归属到分配它的 Python 调用栈。XProf 展示的是剖析窗口内的行为而快照展示的是你选定时刻的设备驻留情况不同时刻的快照还可以做差diff非常适合追踪内存泄漏——即被 Python 引用持有、跨步累积的 buffer。5.1 安装 pprofJAX 设备内存剖析器输出 pprof 格式的数据需要 pprof 工具来解读。安装步骤为先安装 Go 1.16 与 Graphviz然后运行go install github.com/google/pproflatest这会以$GOPATH/bin/pprof安装 pprofGOPATH默认为~/go。注意这里的 pprof 与gperftools包里那个同名老工具不是同一个东西gperftools 版本无法用于 JAX。5.2 理解 JAX 程序如何使用 GPU/TPU 内存设备内存剖析器最常见的用途是搞清 JAX 程序为何占用大量 GPU/TPU 内存例如排查 OOM。用jax.profiler.save_device_memory_profile把设备内存 profile 保存到磁盘。例如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 --http: memory.profpprof 会在浏览器中打开如下 callgraph 形式的设备内存 profile 可视化这个 callgraph 可视化的是每个存活 buffer 被分配时点的 Python 调用栈。例如上例中可视化显示func2及其被调用者负责分配 76.30MB其中 38.15MB 是在从func2到func1的调用内部分配的。callgraph 的解读方法详见 pprof 文档。两个重要事实用jax.jit编译的函数对设备内存剖析器是不透明的jit 编译函数内分配的任何内存都会被整体归属到该函数名下。这也是原文档示例刻意不使用jit的原因。示例中的block_until_ready()是为了确保func2在采集内存 profile 之前完成异步派发的缘故。从实现看device_memory_profilejax/_src/profiler.py通过client.heap_profile()采集堆快照并 gzip 压缩成 pprof 格式的二进制协议缓冲save_device_memory_profile第 453 行只是把它写入文件的便捷封装。剖析机制通过在 JAX 的设备端分配路径上插桩、为每次分配记录 Python 调用栈来工作插桩始终开启device_memory_profile只是提供 API 在任意时刻取快照。5.3 调试内存泄漏在 REPL 里做快速初检时jax.live_arraysjax/_src/api.py 附近会返回后端当前存活的所有数组往往无需任何工具就能发现累积的数组集合。要把内存增长归属到具体代码再用快照。用 JAX 设备内存剖析器追踪内存泄漏可以借助pprof可视化两个不同时刻的设备内存 profile 之间的变化。例如下面的程序把 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()如果只可视化执行结束时的 profilememory9.prof可能看不出循环每轮迭代都在累积更多设备内存分配pprof --http: memory9.profafunction中那个大而固定的分配主导了 profile但它不会随时间增长。对应图见 docs/_static/device_memory_profile_leak1.svg。改用 pprof 的--diff_base功能可视化跨循环迭代的内存变化就能识别内存为何随时间增长pprof --http: --diff_base memory1.prof memory9.prof可视化结果清楚显示内存增长应归属到anotherfunc内部的normal调用。对应图见 docs/_static/device_memory_profile_leak2.svg。这套两个时刻快照做差的方法正是定位跨步内存累积的标准手段。关于 JAX 如何分配 GPU 内存以及 OOM 应对可进一步阅读 docs/201/gpu-memory.md完整的设备内存剖析专题见 docs/device_memory_profiling.md。六、结语度量先行优化在后JAX 的性能工作流可以总结为三层递进先做正确的基准测试记住 JIT 首次编译、异步派发需block_until_ready、32 位默认精度、设备传输开销四个前提再用 XProf 采集计算轨迹定位时间去向识别设备空闲空隙、首次编译、kernel 过碎、集合通信过长等特征最后用设备内存剖析弄清显存占用与泄漏save_device_memory_profile pprof 做差。仓库中 tests/profiler_test.py 提供了start_server、start_trace、TraceAnnotation、远程剖析等 API 的完整端到端用例benchmarks/ 目录下的各类 benchmark 脚本则是本指南技巧的实际运用样例可作为进一步参考。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考