You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

JAX持久化缓存引发错误编译?训练损失异常排查求助

问题描述
  • 现象:注释掉代码中标记为“THIS LINE”的jax.config.update("jax_persistent_cache_min_compile_time_secs", 0)行时,训练损失持续下降;保留该行时,每轮迭代损失恒定,训练无进展。
  • 复现代码:
import jax
import numpy as np
from jax import numpy as jnp
import optax
from tqdm import tqdm
from pathlib import Path
import shutil

cache_path = "/tmp/jax_cache"
if Path(cache_path).exists():
  shutil.rmtree(cache_path)
jax.clear_caches()
jax.config.update("jax_compilation_cache_dir", cache_path)
jax.config.update("jax_persistent_cache_min_compile_time_secs", 0) # THIS LINE

print('initializing...')
params = jnp.array(np.random.normal(size=[10, 11]))

def loss(params):
  return jnp.sum(params ** 4) / 1000.0

print('training...')
opt = optax.lbfgs(0.03)
 
@jax.jit
def do_update(params, opt_state):
  loss_value, params_grad = jax.value_and_grad(loss)(params)
  updates, opt_state = opt.update(
    params_grad,
    opt_state,
    params=params,
    value=loss_value,
    grad=params_grad,
    value_fn=loss,
  )
  params = optax.apply_updates(params, updates)
  return params, opt_state

opt_state = opt.init(params)

def maybe_log(i):
  if i % 100 == 0:
    loss_float = loss(params).item()
    tqdm.write(f'{i} {loss_float}')

for i in tqdm(range(1000)):
  maybe_log(i)
  params, opt_state = do_update(params, opt_state)
  • 环境信息:
    • jax: 0.5.3
    • 硬件: aarch64 CPU
    • python: 3.13.0
问题原因与分析

这个问题并非你的代码错误,而是JAX持久化缓存与Optax LBFGS优化器在特定环境下的兼容性问题,具体原因如下:

  1. 持久化缓存的序列化特性:当设置jax_persistent_cache_min_compile_time_secs=0时,JAX会将所有JIT编译后的函数及其依赖(包括传递给opt.update的value_fn)序列化并持久化到磁盘缓存中。
  2. LBFGS对value_fn的动态依赖:Optax的LBFGS优化器在update过程中,需要调用value_fn执行线搜索调整步长,这一步依赖value_fn能在JIT上下文内动态执行。
  3. aarch64架构下的缓存bug:在你使用的aarch64 CPU环境中,JAX 0.5.3版本的持久化缓存对函数的序列化/反序列化存在问题,导致从缓存加载后的value_fn无法被LBFGS正确调用,线搜索过程失效,最终参数更新量为0,损失保持恒定。
验证与解决思路
  • 验证方式:在x86架构CPU或GPU环境下运行相同代码,大概率不会出现该问题,可佐证是aarch64架构下JAX缓存的兼容性问题。
  • 临时解决方案:
    • 注释掉jax.config.update("jax_persistent_cache_min_compile_time_secs", 0)行,仅使用JAX默认的内存缓存;
    • 升级JAX到0.6.x及以上版本,后续版本修复了部分aarch64架构下的缓存序列化问题;
    • 调整LBFGS调用方式:避免在JIT函数内部传递value_fn,改为预先计算所需损失和梯度,或替换为不需要动态调用value_fn的优化器(如SGD、Adam)。

内容的提问来源于stack exchange,提问作者mwlon

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.12 19:34:57