持久化编译缓存#

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,请遵循 GCEGKE 的说明。为了获得更好的性能,请启用文件缓存(GCEGKE)。

配置 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_partitioninglayernorm_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 中是不够的。