ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

如何用 JAX 持久编译缓存避免重复编译同一批函数

如何用 JAX 持久编译缓存避免重复编译同一批函数 如何用 JAX 持久编译缓存避免重复编译同一批函数【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax当你反复运行同一批jax.jit编译的函数例如训练重启、CI 重复执行、多节点集群冷启动时每个进程重启后都要重新编译一遍相同的计算。JAX 提供了一个可选的磁盘缓存persistent compilation cache启用后JAX 会把编译好的程序副本存到磁盘下次遇到相同或相似的函数时直接命中跳过重编译。本文按“启用缓存 → 配置写入阈值 → 验证命中/排查 miss → 多节点与已知限制”的顺序给出可直接照做的操作路径。前提条件缓存目录必须在第一次编译发生之前设置否则本次运行不会启用如果缓存不在本地文件系统例如放在 GCS需要先安装 etilspip install etils缓存目录被视为可信资源任何能写该目录的用户都能让读缓存的 JAX 进程执行任意代码所以不要把缓存目录放进他人可写的位置见 docs/501/compilation-cache.md 中的警告。启用缓存三种等价的目录设置方式缓存的启用条件就是“缓存位置已设置”。三种方式任选其一即可(1) 环境变量——在 shell 中运行脚本之前设置或写在 Python 脚本顶部export JAX_COMPILATION_CACHE_DIR/tmp/jax_cacheimport os os.environ[JAX_COMPILATION_CACHE_DIR] /tmp/jax_cache(2)jax.config.update()import jax jax.config.update(jax_compilation_cache_dir, /tmp/jax_cache)(3)set_cache_dir()来自 jax/experimental/compilation_cache/compilation_cache.py它内部就是更新jax_compilation_cache_dir实现见 jax/_src/compilation_cache.pyfrom jax.experimental.compilation_cache import compilation_cache as cc cc.set_cache_dir(/tmp/jax_cache)文档给出的最小示例quick start是把目录、阈值一并配好然后正常调用jit函数首次运行即写缓存import jax import jax.numpy as jnp jax.config.update(jax_compilation_cache_dir, /tmp/jax_cache) jax.config.update(jax_persistent_cache_min_entry_size_bytes, -1) jax.config.update(jax_persistent_cache_min_compile_time_secs, 0) jax.config.update(jax_persistent_cache_enable_xla_caches, xla_gpu_per_fusion_autotune_cache_dir) jax.jit def f(x): return x 1 x jnp.zeros((2, 2)) f(x)其中两条阈值配置的含义见下一节——这里置为-1和0是为了让这个小例子也能进缓存。控制哪些函数会被写入缓存两个阈值一个函数要写入持久缓存必须同时满足下面两条标准来自 docs/persistent_compilation_cache.mdjax_persistent_cache_min_compile_time_secs只有编译时间长于该值才会写入默认1.0 秒。可以调高该值来减少写入的条目数jax_persistent_cache_min_entry_size_bytes缓存条目的最小字节数-1关闭大小限制且禁止覆盖默认0允许覆盖——覆盖值通常会把最小尺寸调整到当前缓存文件系统的最优值对应源码中default_min_cache_entry_size()默认返回 0见 jax/_src/compilation_cache.py 0指定你希望的实际最小值且不做覆盖。如果你的函数编译时间很短、条目很小用默认阈值它们会被直接跳过。想让“同一批函数”尽量都落盘比如上面 quick start 那种小函数就需要像示例那样显式把两个阈值调低。另外可以叠加 XLA 自身的缓存机制jax_persistent_cache_enable_xla_caches取值all启用全部 XLA 缓存特性none不启用额外 XLA 缓存xla_gpu_kernel_cache_file只启用 kernel cachexla_gpu_per_fusion_autotune_cache_dir只启用 autotuning cache当前默认值见 jax/_src/config.py 中的default定义。如何确认缓存在工作日志与 cache miss 解释文档建议通过日志观察缓存的实际命中/写入情况。两种方式在脚本顶部设置针对相关模块开启 debug 日志import os os.environ[JAX_DEBUG_LOG_MODULES] jax._src.compiler,jax._src.lru_cache或把全局日志级别调到 DEBUGimport os os.environ[JAX_LOGGING_LEVEL] DEBUG # 或局部使用 jax.config.update(jax_logging_level, DEBUG)如果怀疑某个函数没有命中缓存可以打开 miss 解释源码中默认False每次命中失败时会用logging输出解释jax.config.update(jax_explain_cache_misses, True)注意文档的限定目前该解释仅对 tracing cache 的 miss 实现了说明目标是覆盖所有 miss。所以看到解释日志时先确认它解释的是哪一层缓存。多节点场景缓存放在哪里决定能否全员命中多节点下的行为规则文档“Caching on multiple nodes”一节冷缓存首次运行所有进程都会编译但只有全局通信组中rank 0的进程写缓存后续运行所有进程都尝试读缓存。因此缓存必须放在共享文件系统如 NFS或远程存储如 GCS上如果缓存只本地在 rank 0 上其余进程在后续运行中会因 cache miss 再次编译。如果在 Google Cloud 上文档推荐把缓存放到 GCS bucket并给出了具体配置建议bucket 与 workload 同区域、同项目且 VM 有写权限小 workload 无需复制默认存储类用 “Standard”软删除策略设为最短 7 天用age条件 Delete动作设置对象生命周期覆盖整个 workload 运行期例如预计跑 10 天就设 10 天——因为 JAX 侧没有淘汰机制不设生命周期缓存会持续增长。多节点同时写 GCS 可能触发限流错误因此文档推荐用 GCSFuse 把 bucket 挂载成本地目录GCSFuse 保证同一文件同一时刻只有一个进程写# 假设 GCS bucket 已挂载在 /gcs/my-bucket jax.config.update(jax_compilation_cache_dir, /gcs/my-bucket/jax-cache)不走 GCSFuse 也可以直接指向 bucketjax.config.update(jax_compilation_cache_dir, gs://jax-cache)在单机上为多节点程序预热缓存为了省掉集群上的昂贵编译时间可以在一台机器上用假远程设备把缓存填好。用jax_mock_gpu_topology模拟集群拓扑例如模拟 4 节点、每节点 8 进程、每进程 1 张 GPUjax.config.update(jax_mock_gpu_topology, 4x8x1)用该配置跑一遍程序把缓存填好后就能在真实的 4 节点 × 8 进程 × 1 GPU 拓扑上运行而不再重编译。两条硬性注意运行模拟程序的那台机器其 GPU 数量和 GPU 型号必须与将使用缓存的节点一致。例如8x4x2的模拟拓扑必须在有 2 张 GPU 的机器上运行模拟拓扑下与其他节点的通信结果是未定义的所以模拟环境里 JAX 程序的输出很可能是错误的——它只用于填缓存不用于验证结果。何时缓存会“形同虚设”custom_partitioning 限制一个已知坑使用了自带custom_partitioning的 primitive 的函数缓存不生效。原因是函数 HLO 中包含指向custom_partitioning回调的指针导致同一计算在每次运行产生不同的缓存 key。表现是缓存流程照常进行但每次都生成新 key缓存永远 miss。文档给出的绕行方法用shard_map把那个实现了custom_partitioning的 primitive 包起来。以“LayerNorm 矩阵乘法”为例直接 jit 时每次 key 都不同import jax def F(x1, x2, gamma, beta): ln_out LayerNorm(x1, gamma, beta) return ln_out x2 layernorm_matmul_without_shard_map jax.jit(F, in_shardings(...), out_sharding(...))(x1, x2, gamma, beta)把LayerNormprimitive 包进shard_map后同一计算每次得到相同的缓存 keyPartitionSpec/Mesh的具体参数需按你的实际分片配置替换import jax def G(x1, x2, gamma, beta, mesh, ispecs, ospecs): ln_out jax.shard_map(LayerNorm, meshmesh, in_specsispecs, out_specsospecs, check_vmaFalse)(x1, x2, gamma, beta) return ln_out x2 ispecs jax.sharding.PartitionSpec(...) ospecs jax.sharding.PartitionSpec(...) mesh jax.sharding.Mesh(...) layernorm_matmul_with_shard_map jax.jit(G, static_argnames[mesh, ispecs, ospecs])(x1, x2, gamma, beta, mesh, ispecs, ospecs)注意被shard_map包住的是实现custom_partitioning的那个 primitive只包外层函数F是无效的。出现 miss 时如何判断原因缓存 key 由什么构成排查“为什么没命中”时先看 key 的组成文档“How it works”一节。缓存 key 是编译函数的签名包含函数执行的计算由被哈希的 JAX 函数对应的非优化 HLO 捕获jaxlib 版本相关的 XLA 编译 flags设备配置设备数量与拓扑。目前对 GPU拓扑只包含 GPU 名字符串表示用于压缩编译产物的压缩算法jax._src.cache_key.custom_hook()产生的字符串。该函数可被重指派为用户自定义函数以改变结果字符串默认恒返回空串。也就是说换 jaxlib 版本、换设备/拓扑、换 XLA 编译 flags都会改变 key从而让原本“相同”的函数 miss。如果还发现同一份代码每次 key 都不同优先检查是否踩到了上一节的custom_partitioning问题。限制小结目录未设置或设置晚于首次编译该次运行不启用缓存非本地文件系统需先pip install etils两条阈值编译时间、条目大小必须同时满足默认 1.0 秒的编译时间门槛会跳过快速编译的函数多节点场景缓存必须共享且 JAX 侧没有淘汰机制——放在 GCS 时要自己设置对象生命周期模拟拓扑预热得到的程序输出不可信只用于填缓存使用custom_partitioningprimitive 的函数需要shard_map包裹该 primitive 才能让缓存生效。完整说明可对照 docs/501/compilation-cache.md 与 docs/persistent_compilation_cache.md两篇内容一致后者带 Sphinx 锚点persistent-compilation-cache。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表