使用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缓存的可变状态。 - 细化泄漏定位:
- 在模型
__call__方法中用jax.debug.print打印关键变量的类型,定位出现BatchTracer的位置; - 使用Python的
gc.get_referrers()手动追踪泄漏对象的引用链,确认具体是模型的哪个组件持有了tracer。
- 在模型
- 提供最小复现示例(MWE):若上述步骤无法定位,精简
model代码到能复现泄漏的最小版本,便于进一步排查模型内部的JAX追踪兼容性问题。
内容的提问来源于stack exchange,提问作者thmo
相关产品推荐
相关产品推荐

