JAX 性能回归调查指南:用 nightly 容器与逐小时构建定位引发回归的提交
JAX 性能回归调查指南用 nightly 容器与逐小时构建定位引发回归的提交【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/gh_mirrors/jax/jax更新 JAX 之后发现程序变慢了这是开发者最常遇到的困扰之一。本文基于 docs/investigating_a_regression.md完整复现 JAX 官方在一次 15% 性能回归issue 17686调查中使用的实战方法先用 NVIDIA JAX-Toolbox nightly 容器做粗粒度定位再在固定容器内逐小时同步检出 JAX 与 XLA 源码进行细粒度定位最终锁定引发回归的具体提交。读完本文你将掌握一套不依赖 git bisect、但保证 JAX 与 XLA 始终兼容的暴力穷举调查流程并理解其背后的版本机制与适用边界。一、调查思路为什么不用 git bisect1.1 问题的特殊性JAX 不是一个独立的 Python 包它的底层执行依赖 XLAjaxlib 中包含 XLA。一次JAX 性能回归可能来自JAX Python 层的改动如 jax/_src 下的 API、编译器接入逻辑XLA 层的改动编译优化、代码生成、内存分配策略等。如果只对 JAX 做 git bisect 而 XLA 版本漂移测试结果既不可复现也无意义。因此调查的核心约束是每一次测试都必须保证 JAX 与 XLA 提交是兼容的。1.2 推荐的三步策略原文档给出的调查策略是一个由粗到细的三步流程nightly 容器粗筛在两个发布版本之间用 nightly 容器按天做暴力测试逐小时复筛在固定容器内保持 XLA 与 JAX 同步按小时重新编译并测试最终验证对候选时间窗内的少量提交做手动验证或者使用 git bisect。这是一个 brute force 方法而非严格二分。它的前提是复现脚本足够快——只要单次测试跑得够快暴力穷举就比二分更省事也更稳健。原文特别指出两点优势始终测试相互兼容的 XLA 与 JAX 提交限制了 XLA 的重复编译次数。二、第一阶段nightly 容器按天粗筛2.1 工具与背景这一步依赖 NVIDIA JAX-Toolbox 发布的 nightly 容器镜像名形如ghcr.io/nvidia/jax:nightly-2023-07-06每个镜像内打包了当天互相匹配的 JAX、jaxlib含 XLA以及 CUDA 运行环境。你不需要关心它们各自的版本号——容器本身就保证了兼容性。2.2 两个脚本自动化需要两个脚本test_runner.sh负责启动容器并串行跑完所有日期和test.sh在容器内安装依赖并运行复现测试。以下是原文档中真实用于 issue 17686 的脚本。test_runner.sh——双层循环覆盖 7、8、9 三个月的每一天for m in 7 8 9; do for d in seq -w 1 30; do docker run -v $PWD:/dir --gpusall ghcr.io/nvidia/jax:nightly-2023-0${m}-${d} /bin/bash /dir/test.sh OUT-0${m}-${d} done done要点说明-v $PWD:/dir把当前目录挂载进容器使容器内的/dir/test.sh能访问宿主机脚本--gpusall把宿主 GPU 全部暴露给容器复现是 GPU 性能问题每个镜像的输出重定向到独立的OUT-0${m}-${d}文件便于后续 grep 汇总注意循环里m从 7 到 9日期seq -w保证两位补零与 nightly 镜像命名nightly-2023-0${m}-${d}严格对应。test.sh——在容器内安装缺失依赖并运行基准pip install jmp pyvista numpy matplotlib Rtree trimesh jmp termcolor orbax git clone https://github.com/Autodesk/XLB cd XLB export PYTHONPATH. export CUDA_VISIBLE_DEVICES0 # only 1 GPU is needed python3 examples/performance/MLUPS3d.py 256 200要点说明基准程序是 Autodesk XLB 格子玻尔兹曼求解器中的MLUPS3d.py256×256×256 网格、200 步迭代输出指标为MLUPSMillion Lattice Updates Per Second每秒百万格点更新数数值越高性能越好CUDA_VISIBLE_DEVICES0限定单 GPU排除多卡调度噪声nightly 镜像只保证 JAX 全家桶其他第三方依赖可视化、网格、I/O 相关库需现场安装。2.3 汇总 grep 结果并定位坏日子跑完所有容器后用一条命令汇总全部输出grep MLUPS OUT*这是原文档调查 issue 17686 时得到的真实结果节选关键区间OUT-07-06:MLUPS: 587.9240990200157 OUT-07-07:MLUPS: 587.8907972116419 OUT-07-08:MLUPS: 587.3186499464459 OUT-07-09:MLUPS: 587.3130127722537 OUT-07-10:MLUPS: 587.8526619429658 OUT-07-17:MLUPS: 570.1631097290182 OUT-07-18:MLUPS: 570.2819775617064 OUT-07-19:MLUPS: 570.1672213357352 OUT-07-20:MLUPS: 587.437153685251 OUT-07-21:MLUPS: 587.6702557143142 OUT-07-25:MLUPS: 577.3063618431178 OUT-07-26:MLUPS: 577.2362978080912 OUT-07-27:MLUPS: 577.2101850145785 OUT-07-28:MLUPS: 577.0716349809895 OUT-07-29:MLUPS: 577.4223280707176 OUT-07-30:MLUPS: 577.2255967221336 OUT-08-01:MLUPS: 577.277685388252 OUT-08-02:MLUPS: 577.0137874289354 OUT-08-03:MLUPS: 577.1333281553946 OUT-08-04:MLUPS: 577.305012020407 OUT-08-05:MLUPS: 577.2143988866626 OUT-08-06:MLUPS: 577.2409145495443 OUT-08-07:MLUPS: 577.2602819927345 OUT-08-08:MLUPS: 577.2823738293221 OUT-08-09:MLUPS: 577.3453199728248 OUT-08-11:MLUPS: 577.3161423260563 OUT-08-12:MLUPS: 577.1697775786824 OUT-08-13:MLUPS: 577.3049883393633 OUT-08-14:MLUPS: 576.9051978525331 OUT-08-15:MLUPS: 577.5331743016213 OUT-08-16:MLUPS: 577.5117505070573 OUT-08-18:MLUPS: 577.5930698237612 OUT-08-19:MLUPS: 577.3539885757353 OUT-08-20:MLUPS: 577.4190113959127 OUT-08-21:MLUPS: 577.300394253605 OUT-08-22:MLUPS: 577.4263792037783 OUT-08-23:MLUPS: 577.4087536357031 OUT-08-24:MLUPS: 577.1094728438082 OUT-08-25: File /XLB/examples/performance/MLUPS3d.py, line 5, in module OUT-08-26:MLUPS: 537.0164618489928 OUT-08-27:MLUPS: 536.9545448661609 OUT-08-28:MLUPS: 536.2887650464874 OUT-08-29:MLUPS: 536.7178471720636 OUT-08-30:MLUPS: 536.6978912984252 OUT-09-01:MLUPS: 536.7030899164106 OUT-09-04:MLUPS: 536.5339818238837 OUT-09-05:MLUPS: 536.6507808565617 OUT-09-06:MLUPS: 536.7144494518315 OUT-09-08:MLUPS: 536.7376612408998 OUT-09-09:MLUPS: 536.7798324141778 OUT-09-10:MLUPS: 536.726157440174 OUT-09-11:MLUPS: 536.7446210750584 OUT-09-12:MLUPS: 536.6707332269023 OUT-09-13:MLUPS: 536.6777936517823 OUT-09-14:MLUPS: 536.7581523280307 OUT-09-15:MLUPS: 536.6156273667873 OUT-09-16:MLUPS: 536.7320935035265 OUT-09-17:MLUPS: 536.7104991444398 OUT-09-18:MLUPS: 536.7492269469092 OUT-09-19:MLUPS: 536.6760131792959 OUT-09-20:MLUPS: 536.73612600766342.4 结果解读8 月 24 日 MLUPS ≈ 577属于正常水平8 月 26 日起骤降到 ≈ 537下降了约 7%8 月 25 日的输出是一段 Python 报错MLUPS3d.py第 5 行的异常说明当天容器存在其他构建/运行问题拿不到有效数据结论回归发生在8-24 与 8-26 之间需要把时间窗缩小到这两天之间做逐小时调查。原文还提醒两个实用细节部分日期的镜像可能因构建 bug 失败或本身携带临时回归——直接丢弃这些日期即可不要影响整体判断从完整数据可以看到 7 月中旬还出现过一次更早的轻微性能下降约 570 区间本次调查先忽略它如果需要只要对那段日期再跑一轮 hourly 调查即可。三、第二阶段固定容器内逐小时定位3.1 思路按天的粒度太粗需要把时间窗缩小到小时级。Hourly 调查的做法是以好日子8-24的 nightly 容器为常驻工作容器启动后不退出在容器内对/opt/xla-source与/opt/jax-source两个 git 仓库执行git remote update保持与远端同步对时间窗内的每个小时用git rev-list -1 --before...检出该时刻之前的最新提交重新构建并运行测试。这样做的好处是除第一次构建外之后每次都只做XLA 的增量编译速度远快于反复启停容器。3.2 test_runner2.sh循环驱动# Execute this script inside the container: # docker run -v $PWD:/dir --gpusall ghcr.io/nvidia/jax:nightly-2023-08-24 /bin/bash cd /opt/xla-source git remote update cd /opt/jax-source git remote update pip install jmp pyvista numpy matplotlib Rtree trimesh jmp termcolor orbax cd /tmp git clone https://github.com/Autodesk/XLB cd XLB for d in seq -w 24 26; do for h in seq -w 0 24; do echo $m $d $h /bin/bash /dir/test2.sh Aug $d 2023 $h:00:00 OUT-08-${d}-$h done done要点说明先以docker run ... ghcr.io/nvidia/jax:nightly-2023-08-24 /bin/bash交互式进入容器挂载当前目录到/dir脚本在容器内部执行覆盖 8-24 到 8-26 三天、每天 024 时共 75 个时间点每次调用test2.sh时传入形如Aug 24 2023 12:00:00的时间戳参数输出重定向到OUT-08-24-12这类文件。3.3 test2.sh检出 构建 测试echo param: $ cd /opt/xla-source git checkout git rev-list -1 --before$* origin/main git show -q cd /opt/jax-source git checkout git rev-list -1 --before$* origin/main git show -q rm /opt/jax-source/dist/jax*.whl build-jax.sh # The script is in the nightly container export PYTHONPATH. export CUDA_VISIBLE_DEVICES0 # only 1 GPU is needed python3 examples/performance/MLUPS3d.py 256 200这里有两个值得展开的技术点git rev-list -1 --beforeAug 24 2023 12:00:00 origin/main返回指定时刻之前origin/main上的最新一次提交的哈希随后git checkout精确检出该提交。这是把时间戳翻译成提交的关键命令XLA 和 JAX 两个仓库都这样做从而保证任意时刻两者都取自各自的origin/main、天然同源同步build-jax.sh是 nightly 容器内置的构建脚本。它把 XLA 的增量改动编译进 jaxlib 并生成 wheelrm /opt/jax-source/dist/jax*.whl先清掉旧 wheel避免装错。第一次构建较慢之后每次只重编 XLA 改动因此整个循环是慢启动、快迭代。3.4 汇总对新的输出文件再次执行grep MLUPS OUT-08-*就能看到回归具体出现在哪个小时之间。得到小时级窗口后进入最终验证。四、第三阶段最终验证在小时级窗口内JAX 与 XLA 的历史上通常只有少量提交需要确认查看该时间窗内XLA 与 JAX 各自的提交历史git log --oneline配合时间过滤列出候选提交对每个候选提交手动跑一次复现测试确认哪一次引入回归如果想更 fancy可以在两个已知好坏提交之间用git bisect做二分定位。需要说明这里的 git bisect 与开头的暴力穷举并不冲突——暴力法用于把搜索空间从两个发布版本压缩到几个小时bisect 只在极小窗口内可选使用。五、局限与改进方向原文坦率地指出了这套方法的两点边界如果回归是崩溃crash而不是性能下降那么逐小时重编译跑测试的代价会高得多此时如果能做成真正可自动判定的 bisect 流程会更有用但实现更复杂需要自动判定好/坏对性能回归做二分会丢失信息本次调查的完整数据里其实隐藏着两次回归7 月中旬一次、8 月底一次如果只做 bisect可能只找到其中一次。暴力枚举全部数据点反而更容易看到不止一个回归的全貌。这也是本文方法相对 git bisect 的核心价值以时间顺序穷举全部候选牺牲效率换取信息完整性。六、仓库佐证版本机制与本地构建6.1 nightly 版本号从哪来为什么 nightly 镜像能代表某一天的 JAX从 jax/version.py 可以验证版本生成机制_version 0.4.31为基线版本_get_version_for_build()在设置JAX_NIGHTLY或JAXLIB_NIGHTLY环境变量时会生成形如0.4.31.dev20230906的日期版本从 git 树构建时则追加提交哈希形如0.4.31.dev20230906ge58560fdcg前缀 短哈希_minimum_jaxlib_version 0.4.30表明 JAX 对 jaxlib 有最低版本约束jaxlib 版本必须与 jax 匹配——这正是每次测试必须使用兼容的 XLA/JAX 组合这一核心原则在代码层的体现也解释了为什么调查要整体切换 nightly 容器而不是单独升级 jax。6.2 本地复现与源码构建如果需要在本地复现候选提交而不使用容器可以参考 docs/developer.md 的源码构建流程git clone https://github.com/google/jax cd jax python build/build.py # 构建 jaxlib含 XLA pip install dist/*.whl # 安装 jaxlib 与 jax构建 jaxlib 需要 C 编译器Ubuntu/Debian 可sudo apt install g python python3-devCUDA 版本使用python build/build.py --enable_cuda若需使用本地修改过的 XLA 树可用--bazel_options--override_repositoryxla/path/to/xla让 Bazel 覆盖默认的固定 XLA 版本——这与 hourly 调查中同步检出 XLA 提交的思路一致。6.3 基准程序与性能口径本案例的复现基准XLB 的MLUPS3d.py不在本仓库内但 JAX 仓库自身在 benchmarks 目录下维护了一批官方基准脚本覆盖 API、线性代数、数学函数、随机数、形状多态、稀疏等主题如 benchmarks/api_benchmark.py、benchmarks/linalg_benchmark.py并大量使用jax.block_until_ready()确保测量的是真实计算时间而非异步调度时间。调查你自己的回归时可以参考这些基准的写法把复现脚本固定为稳定输入 固定迭代 单指标输出的形式便于 grep 汇总。6.4 文档在开发者指南中的定位本文档在仓库中位于开发者文档体系内docs/contributor_guide.rst 的 toctree 将其与contributing、developer、jax_internal_api、autodidax、jep/index并列列出说明它是 JAX 社区面向贡献者与资深用户的回归调查方法论类文档。七、小结一份可复用的调查清单最后把整个流程压缩成一张可执行的清单确认回归更新前后分别跑复现基准固定单 GPUCUDA_VISIBLE_DEVICES0、固定输入规模、输出单一指标粗筛按天用 NVIDIA JAX-Toolbox nightly 容器遍历两个版本间的所有日期grep汇总各日指标丢弃构建失败的日期锁定好日子→坏日子边界复筛按小时以好日子的容器为常驻环境用git rev-list -1 --before同步检出 XLA 与 JAX 每小时的提交清空旧 wheel 后增量构建再测试最终验证在小时窗口内审阅两仓库的提交历史对少量候选提交手动确认必要时在极小窗口内使用 git bisect记录结论在向 JAX 提交 issue 时附上引发回归的 commit这将极大加速维护者的定位与修复。这套方法不依赖任何私有工具核心依赖只有一个快速的复现脚本它牺牲了二分效率却换来了不会漏掉多次回归、始终测试兼容版本的稳健性是处理 JAX/XLA 双仓库性能回归时值得优先采用的调查范式。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/gh_mirrors/jax/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考