持久化编译缓存#
JAX 为已编译程序提供了一个可选的磁盘缓存。如果启用,JAX 会将已编译程序的副本存储在磁盘上,这可以在重复运行相同或类似任务时节省重新编译的时间。
注意:如果编译缓存不在本地文件系统上,则需要安装 etils。
pip install etils
用法#
快速入门#
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) 使用环境变量
在运行脚本前的 shell 中:
export JAX_COMPILATION_CACHE_DIR="/tmp/jax_cache"
或者在 Python 脚本的顶部:
import os
os.environ["JAX_COMPILATION_CACHE_DIR"] = "/tmp/jax_cache"
(2) 使用 jax.config.update()
jax.config.update("jax_compilation_cache_dir", "/tmp/jax_cache")
(3) 使用 set_cache_dir()
from jax.experimental.compilation_cache import compilation_cache as cc
cc.set_cache_dir("/tmp/jax_cache")
缓存阈值#
jax_persistent_cache_min_compile_time_secs:仅当编译时间超过指定值时,计算结果才会被写入持久化缓存。默认值为 1.0 秒。jax_persistent_cache_min_entry_size_bytes:将在持久化编译缓存中缓存的条目的最小大小(以字节为单位)。-1:禁用大小限制并防止覆盖。保留默认值(
0)以允许覆盖。覆盖通常会确保最小大小对于所使用的文件系统而言是最佳的。> 0:所需的实际最小大小;无覆盖。
请注意,函数要被缓存,必须同时满足这两个条件。
额外缓存#
XLA 支持额外的缓存机制,可以与 JAX 的持久化编译缓存一起启用,以进一步缩短重新编译时间。
jax_persistent_cache_enable_xla_caches:可选值all:启用所有 XLA 缓存功能none:不启用任何额外的 XLA 缓存功能xla_gpu_kernel_cache_file:仅启用内核缓存xla_gpu_per_fusion_autotune_cache_dir:(默认值)仅启用自动调优缓存
Google Cloud#
在 Google Cloud 上运行时,可以将编译缓存放置在 Google Cloud Storage (GCS) 存储桶中。我们建议采用以下配置:
在工作负载运行所在的同一区域创建存储桶。
在工作负载虚拟机 (VM) 所在的同一项目中创建存储桶。确保设置了权限,以便虚拟机可以写入存储桶。
对于较小的工作负载,无需复制。较大的工作负载可能会从复制中受益。
存储桶的默认存储类别使用“标准”(Standard)。
将软删除策略设置为最短期限:7 天。
将对象生命周期设置为工作负载运行的预期持续时间。例如,如果工作负载预计运行 10 天,请将对象生命周期设置为 10 天。这应该能覆盖整个运行期间发生的重启。使用
age作为生命周期条件,并使用Delete作为操作。详情请参阅 对象生命周期管理。如果未设置对象生命周期,缓存将持续增长,因为没有实现驱逐机制。支持所有加密策略。
建议使用 Google Cloud Storage Fuse 将 GCS 存储桶挂载为本地目录。这是因为在多节点设置中运行 JAX 时,多个节点可能会尝试同时写入缓存,从而导致 GCS 速率限制错误。GCSFuse 通过确保一次只有一个进程可以写入文件来处理此问题,从而防止这些错误。
要设置 GCSFuse,请遵循 GCE 或 GKE 的说明。为了获得更好的性能,请启用文件缓存(GCE 和 GKE)。
配置 GCSFuse 后,将 JAX 缓存目录设置为 GCSFuse 挂载点:
# Example assuming the GCS bucket is mounted at /gcs/my-bucket
jax.config.update("jax_compilation_cache_dir", "/gcs/my-bucket/jax-cache")
直接 GCS 访问
如果您选择不使用 GCSFuse,则可以将缓存直接指向 GCS 存储桶。
假设 gs://jax-cache 是 GCS 存储桶,请按如下方式设置缓存位置:
jax.config.update("jax_compilation_cache_dir", "gs://jax-cache")
工作原理#
缓存键是已编译函数的签名,包含以下参数:
函数执行的计算,由被哈希的 JAX 函数的未优化 HLO 捕获。
jaxlib 版本。
相关的 XLA 编译标志。
通常捕获的设备配置,按设备数量和设备拓扑划分。目前对于 GPU,拓扑仅包含 GPU 名称的字符串表示。
用于压缩已编译可执行文件的压缩算法。
由
jax._src.cache_key.custom_hook()生成的字符串。此函数可以重新分配给用户定义的函数,以便更改生成的字符串。默认情况下,此函数始终返回空字符串。
多节点缓存#
第一次运行程序时(持久化缓存是冷的/空的),所有进程都将进行编译,但只有全局通信组中等级为 0 的进程会将结果写入持久化缓存。在后续运行中,所有进程都会尝试从持久化缓存中读取,因此持久化缓存必须位于共享文件系统(例如 NFS)或远程存储(例如 GFS)中,这一点很重要。如果持久化缓存仅对等级 0 可见,那么在后续运行中,除等级 0 外的所有进程由于编译缓存未命中,都将再次进行编译。
在单节点上预编译多节点程序#
JAX 可以在单节点上为多个节点预填充已编译程序的编译缓存。在单节点上准备缓存有助于减少集群上昂贵的编译时间。要在单节点上编译和运行多节点程序,用户可以使用 jax_mock_gpu_topology 配置选项创建伪远程设备。
例如,下面的代码片段指示 JAX 模拟一个包含四个节点的集群,每个节点运行八个进程,每个进程连接到一个 GPU。
jax.config.update("jax_mock_gpu_topology", "4x8x1")
使用此配置填充缓存后,用户无需重新编译即可在四个节点上运行程序,每个节点八个进程,每个进程一个 GPU。
重要提示
运行模拟程序的进程必须与将使用缓存的节点具有相同数量的 GPU 和相同的 GPU 模型。例如,模拟的拓扑
8x4x2必须在具有两个 GPU 的进程中运行。当运行具有模拟拓扑的程序时,与其他节点的通信结果是未定义的,因此在模拟环境中运行的 JAX 程序的输出很可能是错误的。
记录缓存活动#
检查持久化编译缓存中到底发生了什么对于调试很有帮助。以下是一些入门建议。
用户可以通过放置以下代码来启用相关源文件的日志记录:
import os
os.environ["JAX_DEBUG_LOG_MODULES"] = "jax._src.compiler,jax._src.lru_cache"
在脚本顶部。或者,您可以使用以下命令更改全局 JAX 日志记录级别:
import os
os.environ["JAX_LOGGING_LEVEL"] = "DEBUG"
# or locally with
jax.config.update("jax_logging_level", "DEBUG")
检查缓存未命中#
为了检查和理解缓存未命中的原因,JAX 包含一个配置标志,可以记录所有缓存未命中(包括持久化编译缓存未命中)及其解释。虽然目前这仅针对跟踪缓存未命中实现,但最终目标是解释所有缓存未命中。可以通过设置以下配置来启用此功能。
jax.config.update("jax_explain_cache_misses", True)
潜在缺陷#
目前已经发现了一些潜在缺陷:
目前,持久化缓存不适用于具有主机回调 (host callbacks) 的函数。在这种情况下,完全避免缓存。
这是因为 HLO 包含指向回调的指针,即使计算和计算基础架构完全相同,每次运行也会发生变化。
目前,持久化缓存不适用于使用实现其自身
custom_partitioning的原语 (primitives) 的函数。该函数的 HLO 包含指向
custom_partitioning回调的指针,导致在多次运行中,即使是相同的计算,缓存键也会不同。在这种情况下,缓存仍然会进行,但每次都会产生不同的键,从而使缓存无效。
绕过 custom_partitioning#
如前所述,编译缓存不适用于由实现 custom_partitioning 的原语组成的函数。但是,可以使用 shard_map 为那些实现了 custom_partitioning 的原语绕过该机制,并使编译缓存按预期工作。
假设我们有一个函数 F,它执行层归一化 (layernorm) 后接一个矩阵乘法,其中使用了一个实现 custom_partitioning 的原语 LayerNorm。
import jax
def F(x1, x2, gamma, beta):
ln_out = LayerNorm(x1, gamma, beta)
return ln_out @ x2
如果我们仅仅在没有 shard_map 的情况下编译此函数,layernorm_matmul_without_shard_map 的缓存键在每次运行相同代码时都会不同。
layernorm_matmul_without_shard_map = jax.jit(F, in_shardings=(...), out_sharding=(...))(x1, x2, gamma, beta)
但是,如果我们用 shard_map 包装层归一化原语并定义一个执行相同计算的函数 G,那么尽管 LayerNorm 实现了 custom_partitioning,layernorm_matmul_with_shard_map 的缓存键在每次运行时都将相同。
import jax
def G(x1, x2, gamma, beta, mesh, ispecs, ospecs):
ln_out = jax.shard_map(LayerNorm, mesh=mesh, in_specs=ispecs, out_specs=ospecs, check_vma=False)(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)
请注意,为了实现此变通方法,必须将实现 custom_partitioning 的原语包装在 shard_map 中。仅将外部函数 F 包装在 shard_map 中是不够的。