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

使用vmap时出现JAX Tracer Leak问题,求调试指引

JAX vmap + Equinox梯度计算时Tracer泄漏问题调试指引

问题背景

使用jax.vmap结合eqx.filter_value_and_grad计算批量梯度时,触发JAX Tracer泄漏。单样本梯度函数运行正常,但批量版本在jax.check_tracer_leaks()检测下报错,泄漏路径为:BatchTracer→列表→PjitParams→InferParamsCacheEntry。

梯度函数定义

@eqx.filter_value_and_grad
def grad_loss(model, ti, yi):
    y_pred = model(ti, yi[0])
    return jnp.mean((yi - y_pred) ** 2)

@eqx.filter_value_and_grad
def grad_loss_batch(model, ti, yi):
    y_pred = jax.vmap(model, (None, 0))(ti, yi[:, 0])
    return jnp.mean((yi - y_pred) ** 2)

测试代码与报错

ts = jnp.linspace(0.0, 1.0, 10)
ys = jax.lax.stop_gradient(jnp.ones_like(ts)[..., None])

# 单样本无泄漏
with jax.checking_leaks():
   loss, grads = grad_loss(model_dde, ts, ys)

# 批量版本触发泄漏
with jax.check_tracer_leaks():
    loss2, grads2 = grad_loss_batch(model_dde, ts, ys[None, ...])

报错信息:

*** Exception: Leaked trace MainTrace(3,BatchTrace). Leaked tracer(s):

Traced<ShapedArray(float32[1])>with<BatchTrace(level=3/0)> with
  val = Array([[1.]], dtype=float32)
  batch_dim = 0
This BatchTracer with object id 132573553507072 was created on line:
  /home/monsel/Desktop/dev_diffrax/mwe.py:69 (grad_loss_batch)
<BatchTracer 132573553507072> is referred to by <list 132573501535936>[0]
<list 132573501535936> is referred to by <PjitParams 132573554102464>[0]
<PjitParams 132573554102464> is referred to by <InferParamsCacheEntry 132573553590720>

调试步骤

  • 排查模型内部的可变状态/缓存:泄漏路径指向JIT缓存相关结构,说明model的__call__方法可能存在副作用(比如修改实例属性、向列表/字典添加元素),导致vmap追踪时的BatchTracer被意外存储。检查模型代码中是否有这类操作,确保模型调用是纯函数(无状态修改)。
  • 验证模型纯函数性:用极简纯函数模型(如线性变换)替换model_dde,测试批量梯度函数是否还泄漏。如果问题消失,可确定是原模型的状态依赖导致的。
  • 调整vmap与梯度装饰器的嵌套顺序:将vmap移到梯度函数外面,对单样本梯度函数做批量处理,而非在梯度函数内部vmap模型调用:
    @eqx.filter_value_and_grad
    def grad_loss(model, ti, yi):
        y_pred = model(ti, yi[0])
        return jnp.mean((yi - y_pred) ** 2)
    
    # 对单样本梯度函数做vmap
    grad_loss_batch = jax.vmap(grad_loss, in_axes=(None, None, 0))
    
  • 禁用JIT验证:设置环境变量JAX_DISABLE_JIT=1运行测试,若泄漏消失,说明是JIT缓存捕获了tracer。此时需检查模型中是否有被JIT缓存的可变状态。
  • 细化泄漏定位:
    1. 在模型__call__方法中用jax.debug.print打印关键变量的类型,定位出现BatchTracer的位置;
    2. 使用Python的gc.get_referrers()手动追踪泄漏对象的引用链,确认具体是模型的哪个组件持有了tracer。
  • 提供最小复现示例(MWE):若上述步骤无法定位,精简model代码到能复现泄漏的最小版本,便于进一步排查模型内部的JAX追踪兼容性问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 19:07:33